Spaces:
Sleeping
Sleeping
File size: 5,345 Bytes
2415446 61bb677 2415446 7449374 2415446 7449374 61bb677 2415446 7449374 2415446 7449374 2415446 7449374 61bb677 2415446 7449374 2415446 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 | """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,
reasoning=routed.reasoning,
)
gateway_model = routed.resolved.original_model
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": gateway_model,
"reasoning_control": routed.reasoning.control.value,
"reasoning_effort": (
routed.reasoning.effort.value
if routed.reasoning.effort is not None
else None
),
"reasoning_budget_tokens": routed.reasoning.budget_tokens,
}
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)
request_snapshot = anthropic_request_snapshot(routed.request)
request_snapshot["model"] = gateway_model
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=request_snapshot,
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,
response_model=gateway_model,
reasoning=routed.reasoning,
)
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": gateway_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,
)
|