refactor(agent/codex_runtime,bedrock_adapter): fold assembler state, pending-call guard, toolUse block builder (-16 LOC)

This commit is contained in:
Teknium
2026-09-02 19:26:47 -07:00
parent db1214a406
commit 72f1e2e2c1
2 changed files with 31 additions and 47 deletions

View File

@@ -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, "")

View File

@@ -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.