refactor(agent/codex_runtime,bedrock_adapter): fold assembler state, pending-call guard, toolUse block builder (-16 LOC)
This commit is contained in:
@@ -550,6 +550,10 @@ def _tool_use_block(tool_use_id, name, input_dict) -> Dict:
|
||||
return {"toolUse": {"toolUseId": tool_use_id, "name": name, "input": input_dict}}
|
||||
|
||||
|
||||
def _tool_use_block_from(tu: Dict) -> Dict:
|
||||
return _tool_use_block(tu.get("toolUseId", ""), tu.get("name", ""), tu.get("input", {}))
|
||||
|
||||
|
||||
def _decode_redacted(encoded) -> Optional[bytes]:
|
||||
"""Strict base64 → bytes; None for empty/non-str/undecodable input."""
|
||||
if not isinstance(encoded, str) or not encoded:
|
||||
@@ -585,8 +589,7 @@ def _replay_ordered_blocks(ordered_blocks: List) -> List[Dict]:
|
||||
if replay:
|
||||
content_blocks.append({"reasoningContent": replay})
|
||||
elif "toolUse" in block and isinstance(block["toolUse"], dict):
|
||||
tu = block["toolUse"]
|
||||
content_blocks.append(_tool_use_block(tu.get("toolUseId", ""), tu.get("name", ""), tu.get("input", {})))
|
||||
content_blocks.append(_tool_use_block_from(block["toolUse"]))
|
||||
return content_blocks
|
||||
|
||||
|
||||
@@ -712,11 +715,9 @@ class _ResponseParts:
|
||||
"""Assemble the OpenAI-shaped response. Converse's inputTokens EXCLUDES cache
|
||||
read/write tokens (OpenAI's prompt_tokens includes them), so they are added back."""
|
||||
msg = SimpleNamespace(
|
||||
role="assistant",
|
||||
content="\n".join(self.text_parts) if self.text_parts else None,
|
||||
tool_calls=self.tool_calls or None,
|
||||
role="assistant", content="\n".join(self.text_parts) if self.text_parts else None,
|
||||
tool_calls=self.tool_calls or None, reasoning_details=self.reasoning_details or None,
|
||||
reasoning_content="\n\n".join(self.reasoning_parts) if self.reasoning_parts else None,
|
||||
reasoning_details=self.reasoning_details or None,
|
||||
bedrock_content_blocks=ordered_blocks or None,
|
||||
)
|
||||
cache_read_tokens = usage_data.get("cacheReadInputTokens", 0)
|
||||
@@ -724,11 +725,8 @@ class _ResponseParts:
|
||||
output_tokens = usage_data.get("outputTokens", 0)
|
||||
prompt_tokens = usage_data.get("inputTokens", 0) + cache_read_tokens + cache_write_tokens
|
||||
usage = SimpleNamespace(
|
||||
prompt_tokens=prompt_tokens,
|
||||
completion_tokens=output_tokens,
|
||||
total_tokens=prompt_tokens + output_tokens,
|
||||
cache_read_input_tokens=cache_read_tokens,
|
||||
cache_creation_input_tokens=cache_write_tokens,
|
||||
prompt_tokens=prompt_tokens, completion_tokens=output_tokens, total_tokens=prompt_tokens + output_tokens,
|
||||
cache_read_input_tokens=cache_read_tokens, cache_creation_input_tokens=cache_write_tokens,
|
||||
)
|
||||
finish_reason = _STOP_REASON_TO_FINISH_REASON.get(stop_reason, "stop")
|
||||
if self.tool_calls and finish_reason == "stop":
|
||||
@@ -755,9 +753,8 @@ def normalize_converse_response(response: Dict) -> SimpleNamespace:
|
||||
ordered_blocks.append({"reasoningContent": ordered_reasoning})
|
||||
elif "toolUse" in block:
|
||||
tu = block["toolUse"]
|
||||
tool_use_id, name, tool_input = tu.get("toolUseId", ""), tu.get("name", ""), tu.get("input", {})
|
||||
ordered_blocks.append(_tool_use_block(tool_use_id, name, tool_input))
|
||||
parts.tool_calls.append(_tool_call_ns(tool_use_id, name, tool_input))
|
||||
ordered_blocks.append(_tool_use_block_from(tu))
|
||||
parts.tool_calls.append(_tool_call_ns(tu.get("toolUseId", ""), tu.get("name", ""), tu.get("input", {})))
|
||||
return parts.build(
|
||||
ordered_blocks, response.get("usage", {}), response.get("stopReason", "end_turn"), response.get("modelId", ""),
|
||||
)
|
||||
@@ -852,10 +849,7 @@ def stream_converse_with_callbacks(
|
||||
stop_reason = event["messageStop"].get("stopReason", "end_turn")
|
||||
elif "metadata" in event:
|
||||
meta_usage = event["metadata"].get("usage", {})
|
||||
usage_data = {
|
||||
key: meta_usage.get(key, 0)
|
||||
for key in ("inputTokens", "outputTokens", "cacheReadInputTokens", "cacheWriteInputTokens")
|
||||
}
|
||||
usage_data = {key: meta_usage.get(key, 0) for key in ("inputTokens", "outputTokens", "cacheReadInputTokens", "cacheWriteInputTokens")}
|
||||
flush_text()
|
||||
return parts.build([stream_blocks[i] for i in sorted(stream_blocks)], usage_data, stop_reason, "")
|
||||
|
||||
|
||||
@@ -96,8 +96,10 @@ def _record_codex_app_server_usage(agent, turn) -> dict[str, Any]:
|
||||
agent.session_api_calls += 1
|
||||
usage = getattr(turn, "token_usage_last", None)
|
||||
compressor = getattr(agent, "context_compressor", None)
|
||||
|
||||
def billing(**extra):
|
||||
return dict(model=agent.model, billing_provider=agent.provider, billing_base_url=agent.base_url, api_call_count=1, **extra)
|
||||
|
||||
if not isinstance(usage, dict) or not usage:
|
||||
if compressor is not None and getattr(compressor, "awaiting_real_usage_after_compression", False):
|
||||
# No usage cannot adjudicate the pending compaction; unlatch preflight deferral.
|
||||
@@ -109,12 +111,9 @@ def _record_codex_app_server_usage(agent, turn) -> dict[str, Any]:
|
||||
return {}
|
||||
from agent.usage_pricing import CanonicalUsage, estimate_usage_cost
|
||||
canonical_usage = CanonicalUsage(
|
||||
input_tokens=_coerce_usage_int(usage.get("inputTokens")),
|
||||
output_tokens=_coerce_usage_int(usage.get("outputTokens")),
|
||||
cache_read_tokens=_coerce_usage_int(usage.get("cachedInputTokens")),
|
||||
cache_write_tokens=0,
|
||||
reasoning_tokens=_coerce_usage_int(usage.get("reasoningOutputTokens")),
|
||||
raw_usage=usage,
|
||||
input_tokens=_coerce_usage_int(usage.get("inputTokens")), output_tokens=_coerce_usage_int(usage.get("outputTokens")),
|
||||
cache_read_tokens=_coerce_usage_int(usage.get("cachedInputTokens")), cache_write_tokens=0,
|
||||
reasoning_tokens=_coerce_usage_int(usage.get("reasoningOutputTokens")), raw_usage=usage,
|
||||
)
|
||||
prompt_tokens = canonical_usage.prompt_tokens
|
||||
total_tokens = _coerce_usage_int(usage.get("totalTokens")) or canonical_usage.total_tokens
|
||||
@@ -137,8 +136,7 @@ def _record_codex_app_server_usage(agent, turn) -> dict[str, Any]:
|
||||
for key, value in usage_dict.items():
|
||||
setattr(agent, f"session_{key}", getattr(agent, f"session_{key}") + value)
|
||||
cost_result = estimate_usage_cost(
|
||||
agent.model, canonical_usage,
|
||||
provider=agent.provider, base_url=agent.base_url, api_key=getattr(agent, "api_key", ""),
|
||||
agent.model, canonical_usage, provider=agent.provider, base_url=agent.base_url, api_key=getattr(agent, "api_key", ""),
|
||||
)
|
||||
cost_usd = float(cost_result.amount_usd) if cost_result.amount_usd is not None else None
|
||||
if cost_usd is not None:
|
||||
@@ -344,12 +342,11 @@ def make_codex_app_server_event_bridge(agent) -> Callable[[dict], None]:
|
||||
prior = started.pop(item_id, None) if (item_id := item.get("id") or "") else None
|
||||
# Prefer codex's durationMs; else our started timestamp; else None
|
||||
# (some codex versions only emit completed for fast items).
|
||||
duration: Any = None
|
||||
codex_ms = item.get("durationMs")
|
||||
if isinstance(codex_ms, (int, float)) and codex_ms >= 0:
|
||||
duration = codex_ms / 1000.0
|
||||
elif prior is not None:
|
||||
duration = time.monotonic() - prior[2]
|
||||
duration: Any = codex_ms / 1000.0
|
||||
else:
|
||||
duration = time.monotonic() - prior[2] if prior is not None else None
|
||||
result, is_error = _codex_item_completion_payload(item)
|
||||
agent_cb("tool_progress_callback", "tool_progress_callback raised on tool.completed for %s", name,
|
||||
args=("tool.completed", name, None, None),
|
||||
@@ -392,7 +389,7 @@ def make_codex_app_server_event_bridge(agent) -> Callable[[dict], None]:
|
||||
def on_event(note: dict) -> None:
|
||||
handler = handlers.get(note.get("method") or "") if isinstance(note, dict) else None
|
||||
if handler is not None:
|
||||
params = note.get("params") or {}
|
||||
params = note.get("params")
|
||||
handler(params if isinstance(params, dict) else {})
|
||||
|
||||
return on_event
|
||||
@@ -631,20 +628,15 @@ class _CodexResponseAssembler:
|
||||
from text deltas, or settled from function calls announced via ``output_item.added``
|
||||
but never confirmed (some backends omit per-item done events on success)."""
|
||||
|
||||
has_tool_calls = False
|
||||
has_tool_calls = first_delta_fired = saw_terminal = False
|
||||
next_output_sequence = 0
|
||||
first_delta_fired = False
|
||||
active_message_phase: str | None = None
|
||||
# Reasoning summary parts carry no separator; a summary_index change is where the blank line belongs.
|
||||
active_summary_index: Any = None
|
||||
terminal_status: str = "completed"
|
||||
terminal_usage: Any = None
|
||||
terminal_response_id: str = None
|
||||
terminal_incomplete_details: Any = None
|
||||
terminal_error: Any = None
|
||||
saw_terminal = False
|
||||
# terminal_status defaults to "completed", so settlement needs an
|
||||
# explicitly observed response.completed frame (not EOF/interrupt).
|
||||
terminal_usage = terminal_response_id = terminal_incomplete_details = terminal_error = None
|
||||
# terminal_status defaults to "completed", so settlement needs an explicitly
|
||||
# observed response.completed frame (not EOF/interrupt).
|
||||
saw_response_completed = False
|
||||
|
||||
def __init__(self, *, model, on_text_delta, on_reasoning_delta, on_commentary_message, on_first_delta):
|
||||
@@ -717,17 +709,16 @@ class _CodexResponseAssembler:
|
||||
def _on_function_call(self, event: Any, event_type: str) -> None:
|
||||
self.has_tool_calls = True
|
||||
pending = self.pending_function_calls.get(str(_event_field(event, "item_id", "")))
|
||||
if pending is None:
|
||||
return # the item itself lands on output_item.done
|
||||
if "delta" in event_type:
|
||||
delta_args = _event_field(event, "delta", "")
|
||||
if pending is not None and delta_args:
|
||||
pending["arguments"] += delta_args
|
||||
pending["arguments"] += _event_field(event, "delta", "") or ""
|
||||
elif event_type.endswith("function_call_arguments.done"):
|
||||
# Authoritative for the accumulated string; an explicit "" (zero-arg
|
||||
# call) counts, only a missing field keeps the streamed deltas.
|
||||
done_args = _event_field(event, "arguments", None)
|
||||
if pending is not None and done_args is not None:
|
||||
if done_args is not None:
|
||||
pending["arguments"] = str(done_args)
|
||||
# Other function_call frames: the item itself lands on output_item.done.
|
||||
|
||||
def _on_reasoning_delta(self, event: Any, event_type: str) -> None:
|
||||
reasoning_text = _event_field(event, "delta", "")
|
||||
@@ -750,8 +741,7 @@ class _CodexResponseAssembler:
|
||||
done_id = str(_event_field(done_item, "id", ""))
|
||||
announced_sequence, announced_index = self.announced_output_order.get(done_id, (None, None))
|
||||
if announced_sequence is None:
|
||||
announced_sequence = self.next_output_sequence
|
||||
self.next_output_sequence += 1
|
||||
announced_sequence, self.next_output_sequence = self.next_output_sequence, self.next_output_sequence + 1
|
||||
self.output_indexes.append(_event_field(event, "output_index", announced_index))
|
||||
self.output_sequences.append(announced_sequence)
|
||||
# Confirmed by the authoritative done event; never settle it twice.
|
||||
|
||||
Reference in New Issue
Block a user