mirror of
https://github.com/Alishahryar1/free-claude-code.git
synced 2026-07-03 14:05:26 +02:00
112 lines
3.6 KiB
Python
112 lines
3.6 KiB
Python
"""Shared provider execution primitive for API product handlers."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import uuid
|
|
from collections.abc import AsyncIterator, Callable
|
|
from typing import Any
|
|
|
|
from loguru import logger
|
|
|
|
from config.settings import Settings
|
|
from core.anthropic import get_token_count
|
|
from core.trace import api_messages_request_snapshot, trace_event, traced_async_stream
|
|
from providers.base import BaseProvider
|
|
|
|
from .model_router import RoutedMessagesRequest
|
|
|
|
TokenCounter = Callable[[list[Any], str | list[Any] | None, list[Any] | None], int]
|
|
ProviderGetter = Callable[[str], BaseProvider]
|
|
|
|
|
|
class ProviderExecutionService:
|
|
"""Resolve a provider and execute one routed Anthropic Messages stream."""
|
|
|
|
def __init__(
|
|
self,
|
|
settings: Settings,
|
|
provider_getter: ProviderGetter,
|
|
*,
|
|
token_counter: TokenCounter = get_token_count,
|
|
) -> None:
|
|
self._settings = settings
|
|
self._provider_getter = provider_getter
|
|
self._token_counter = token_counter
|
|
|
|
def stream(
|
|
self,
|
|
routed: RoutedMessagesRequest,
|
|
*,
|
|
wire_api: str,
|
|
raw_log_label: str,
|
|
raw_log_payload: Any,
|
|
) -> AsyncIterator[str]:
|
|
provider = self._provider_getter(routed.resolved.provider_id)
|
|
provider.preflight_stream(
|
|
routed.request,
|
|
thinking_enabled=routed.resolved.thinking_enabled,
|
|
)
|
|
|
|
route_trace: dict[str, Any] = {
|
|
"stage": "routing",
|
|
"event": "api.route.resolved",
|
|
"source": "api",
|
|
"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"
|
|
trace_event(**route_trace)
|
|
|
|
request_id = f"req_{uuid.uuid4().hex[:12]}"
|
|
trace_event(
|
|
stage="ingress",
|
|
event=(
|
|
"api.responses.request.received"
|
|
if wire_api == "responses"
|
|
else "api.request.received"
|
|
),
|
|
source="api",
|
|
message_count=len(routed.request.messages),
|
|
snapshot=api_messages_request_snapshot(routed.request),
|
|
request_id=request_id,
|
|
)
|
|
|
|
if self._settings.log_raw_api_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,
|
|
)
|
|
return traced_async_stream(
|
|
provider.stream_response(
|
|
routed.request,
|
|
input_tokens=input_tokens,
|
|
request_id=request_id,
|
|
thinking_enabled=routed.resolved.thinking_enabled,
|
|
),
|
|
stage="egress",
|
|
source="api",
|
|
complete_event=(
|
|
"api.responses.stream_completed"
|
|
if wire_api == "responses"
|
|
else "api.response.stream_completed"
|
|
),
|
|
interrupted_event=(
|
|
"api.responses.stream_interrupted"
|
|
if wire_api == "responses"
|
|
else "api.response.stream_interrupted"
|
|
),
|
|
chunk_event=None,
|
|
extra={
|
|
"request_id": request_id,
|
|
"provider_id": routed.resolved.provider_id,
|
|
"gateway_model": routed.request.model,
|
|
},
|
|
)
|