Files
alishahryar1--free-claude-code/src/free_claude_code/application/execution.py
T
wehub-resource-sync 5296d0e97c
CI / Ban suppressions and legacy annotations (push) Has been cancelled
CI / pytest (push) Has been cancelled
CI / ruff-check (push) Has been cancelled
CI / ruff-format (push) Has been cancelled
CI / ty (push) Has been cancelled
chore: import upstream snapshot with attribution
2026-07-13 12:35:44 +08:00

148 lines
4.8 KiB
Python

"""Provider execution shared by inbound API adapters."""
import sys
from collections.abc import AsyncIterator, Callable
from typing import Literal
from loguru import logger
from free_claude_code.core.anthropic import (
Message,
SystemContent,
Tool,
anthropic_request_snapshot,
get_token_count,
)
from free_claude_code.core.trace import (
close_stream_input,
trace_event,
traced_async_stream,
)
from .ports import ProviderResolver
from .routing import RoutedMessagesRequest
TokenCounter = Callable[
[list[Message], str | list[SystemContent] | None, list[Tool] | None],
int,
]
WireApi = Literal["messages", "responses"]
class ProviderExecutor:
"""Resolve a provider and execute one routed Anthropic Messages stream."""
def __init__(
self,
provider_resolver: ProviderResolver,
*,
token_counter: TokenCounter = get_token_count,
generation_id: int | None = None,
log_raw_payloads: bool = False,
) -> None:
self._provider_resolver = provider_resolver
self._token_counter = token_counter
self._generation_id = generation_id
self._log_raw_payloads = log_raw_payloads
def stream(
self,
routed: RoutedMessagesRequest,
*,
wire_api: WireApi,
raw_log_label: str,
raw_log_payload: object,
request_id: str,
) -> AsyncIterator[str]:
"""Preflight synchronously, then return the traced provider stream."""
provider = self._provider_resolver(routed.resolved.provider_id)
provider.preflight_stream(
routed.request,
thinking_enabled=routed.resolved.thinking_enabled,
)
route_trace: dict[str, object] = {
"stage": "routing",
"event": "free_claude_code.api.route.resolved",
"source": "api",
"request_id": request_id,
"provider_id": routed.resolved.provider_id,
"provider_model": routed.resolved.provider_model,
"provider_model_ref": routed.resolved.provider_model_ref,
"gateway_model": routed.request.model,
"thinking_enabled": routed.resolved.thinking_enabled,
}
if wire_api == "responses":
route_trace["wire_api"] = "responses"
if self._generation_id is not None:
route_trace["generation_id"] = self._generation_id
trace_event(**route_trace)
trace_event(
stage="ingress",
event=(
"free_claude_code.api.responses.request.received"
if wire_api == "responses"
else "free_claude_code.api.request.received"
),
source="api",
message_count=len(routed.request.messages),
snapshot=anthropic_request_snapshot(routed.request),
request_id=request_id,
)
if self._log_raw_payloads:
logger.debug(f"{raw_log_label} [{{}}]: {{}}", request_id, raw_log_payload)
input_tokens = self._token_counter(
routed.request.messages,
routed.request.system,
routed.request.tools,
)
async def provider_body() -> AsyncIterator[str]:
provider_stream: AsyncIterator[str] | None = None
try:
provider_stream = provider.stream_response(
routed.request,
input_tokens=input_tokens,
request_id=request_id,
thinking_enabled=routed.resolved.thinking_enabled,
)
async for chunk in provider_stream:
yield chunk
finally:
if provider_stream is not None:
await close_stream_input(
provider_stream,
owner="provider_executor",
source="api",
preserved_error=sys.exception(),
)
stream_trace: dict[str, object] = {
"request_id": request_id,
"provider_id": routed.resolved.provider_id,
"gateway_model": routed.request.model,
}
if self._generation_id is not None:
stream_trace["generation_id"] = self._generation_id
return traced_async_stream(
provider_body(),
stage="egress",
source="api",
complete_event=(
"free_claude_code.api.responses.stream_completed"
if wire_api == "responses"
else "free_claude_code.api.response.stream_completed"
),
interrupted_event=(
"free_claude_code.api.responses.stream_interrupted"
if wire_api == "responses"
else "free_claude_code.api.response.stream_interrupted"
),
chunk_event=None,
extra=stream_trace,
)