File size: 6,215 Bytes
d29db6e
ce6ad01
4aecccf
e8f2acc
ce6ad01
d29db6e
ce6ad01
 
 
d29db6e
 
 
 
 
 
0d0f0f9
4aecccf
 
 
 
 
e8f2acc
 
 
 
 
 
ce6ad01
d29db6e
 
ce6ad01
d29db6e
 
 
 
 
ce6ad01
 
d29db6e
ce6ad01
 
 
 
d29db6e
ce6ad01
 
0602b26
d29db6e
ce6ad01
d29db6e
ce6ad01
0602b26
d29db6e
ce6ad01
 
 
 
 
d29db6e
ce6ad01
d29db6e
fcd6f91
ce6ad01
d29db6e
 
ce6ad01
 
dc04eee
ce6ad01
 
82e8d6e
d29db6e
ce6ad01
0d0f0f9
ce6ad01
fcd6f91
ce6ad01
 
 
82e8d6e
dc04eee
 
 
 
 
 
 
ce6ad01
 
 
0602b26
 
ce6ad01
 
82e8d6e
 
ce6ad01
 
 
0d0f0f9
ce6ad01
0d0f0f9
ce6ad01
 
 
82e8d6e
ce6ad01
 
 
d29db6e
ce6ad01
 
 
 
 
 
 
64dfa67
e8f2acc
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
64dfa67
4aecccf
 
 
 
 
 
82e8d6e
dc04eee
4aecccf
 
 
 
 
 
 
 
 
 
 
64dfa67
d29db6e
0602b26
 
82e8d6e
0602b26
 
 
 
e8f2acc
 
 
 
64dfa67
e8f2acc
ce6ad01
 
 
0d0f0f9
ce6ad01
0d0f0f9
ce6ad01
 
0d0f0f9
ce6ad01
0d0f0f9
ce6ad01
 
0602b26
ce6ad01
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
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
"""Provider execution shared by inbound API adapters."""

import sys
import time
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 free_claude_code.core.usage_tracking import (
    PendingUsageRecord,
    UsageTrackingStream,
    extract_prompt,
    get_buffer,
)

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

        usage_tracking = get_buffer() is not None
        pending = (
            PendingUsageRecord(
                request_id=request_id,
                started_at=time.time(),
                provider=routed.resolved.provider_id,
                provider_model=routed.resolved.provider_model,
                gateway_model=gateway_model,
                wire_api=wire_api,
                input_tokens=input_tokens,
                prompt=extract_prompt(routed.request),
            )
            if usage_tracking
            else None
        )

        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

        capture_stream: AsyncIterator[str] = provider_body()
        if pending is not None:
            capture_stream = UsageTrackingStream(capture_stream, pending)

        return traced_async_stream(
            capture_stream,
            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,
        )