1"""Streaming pipeline for :class:`RelayGatewayService`.
2
3Preflight (authorization, channel selection, billing admission, request
4conversion, stream-session creation) completes before the first frame is
5delivered; upstream I/O happens lazily as the returned stream is
6consumed by the caller. Billing settles exactly once when the stream
7ends — completed, cancelled, truncated, or failed.
8"""
9
10from __future__ import annotations
11
12import asyncio
13from collections.abc import AsyncGenerator, AsyncIterator
14from typing import Literal, cast
15
16from lexigram.ai.relay.gateway.channels import RelayChannelRegistry
17from lexigram.ai.relay.gateway.codec import RelayPayloadCodec
18from lexigram.ai.relay.gateway.config import RelayGatewayConfig
19from lexigram.ai.relay.gateway.errors import (
20 auth_denied,
21 conversion_error_to_gateway,
22 with_request_id,
23)
24from lexigram.ai.relay.gateway.operations import billing as billing_ops
25from lexigram.ai.relay.gateway.operations import telemetry
26from lexigram.ai.relay.gateway.operations import upstream as upstream_ops
27from lexigram.ai.relay.gateway.operations.failover import RelayFailoverTracker
28from lexigram.ai.relay.gateway.operations.streams import RelayStreamRegistry
29from lexigram.ai.relay.gateway.stream import UpstreamEventParser, relay_stream
30from lexigram.ai.relay.gateway.upstream import HTTPUpstreamAdapter
31from lexigram.contracts.ai.exceptions import RelayError
32from lexigram.contracts.ai.governance import RelayBillingProtocol, RelayUsageReservation
33from lexigram.contracts.ai.relay import (
34 MediaResolverProtocol,
35 RelayChannel,
36 RelayConversionContext,
37 RelayConverterProtocol,
38 RelayConvertResult,
39 RelayGatewayError,
40 RelayGatewayMetadata,
41 RelayGatewayRequest,
42 RelayGatewayResult,
43 RelayOptions,
44 RelayRequestPayload,
45 RelayStreamSessionProtocol,
46 RelayUpstreamProtocol,
47 RelayWireEvent,
48 UpstreamRequest,
49)
50from lexigram.contracts.auth.guard import AuthorizerProtocol
51from lexigram.contracts.core.result import Err, Ok, Result
52from lexigram.logging import get_logger
53
54logger = get_logger(__name__)
55
56
57class StreamingMixin:
58 """Streaming request pipeline shared into ``RelayGatewayService``.
59
60 Requires the host service to provide the gateway dependencies as
61 private attributes (``_converter``, ``_codec``, ``_registry``,
62 ``_upstream``, ``_config``, ``_authorizer``, ``_billing``,
63 ``_media_resolver``, ``_streams``, ``_failover``).
64 """
65
66 _converter: RelayConverterProtocol
67 _codec: RelayPayloadCodec
68 _registry: RelayChannelRegistry
69 _upstream: HTTPUpstreamAdapter
70 _config: RelayGatewayConfig
71 _authorizer: AuthorizerProtocol | None
72 _billing: RelayBillingProtocol | None
73 _media_resolver: MediaResolverProtocol | None
74 _streams: RelayStreamRegistry | None
75 _failover: RelayFailoverTracker | None
76
77 async def _handle_streaming(
78 self, request: RelayGatewayRequest
79 ) -> tuple[Result[RelayGatewayResult, RelayGatewayError], str]:
80 """Run the streaming preflight and return a lazy stream result.
81
82 Authorization, channel selection, billing admission, request
83 conversion, and stream-session creation all complete before the
84 first frame is delivered; upstream I/O happens lazily as the
85 returned stream is consumed by the caller.
86
87 Returns:
88 ``tuple`` of the pipeline result (whose ``stream`` holds the
89 lazy ``AsyncIterator`` on success) and the selected channel
90 name. The channel name is ``""`` when selection failed
91 before a channel was chosen.
92 """
93 if self._authorizer is not None:
94 allowed = await self._authorizer.authorize(
95 user=request.tenant_id,
96 action="relay.invoke",
97 resource=request.model,
98 )
99 if not allowed:
100 return Err(auth_denied(request.request_id)), ""
101 selected = self._registry.select(
102 source=request.source,
103 model=request.model,
104 stream=True,
105 preferred=request.channel.name if request.channel else None,
106 )
107 if selected.is_err():
108 return (
109 Err(with_request_id(selected.unwrap_err(), request.request_id)),
110 "",
111 )
112 channel = selected.unwrap()
113 logger.info(
114 "relay_gateway_channel_selected",
115 request_id=request.request_id,
116 channel=channel.name,
117 target_format=channel.target_format,
118 model=request.model,
119 )
120 billing = self._billing
121 reservation: RelayUsageReservation | None = None
122 if billing is not None:
123 admitted = await billing_ops.pre_consume(
124 self._codec, request, billing, channel
125 )
126 if admitted.is_err():
127 return Err(admitted.unwrap_err()), channel.name
128 reservation = admitted.unwrap()
129 outbound_model = upstream_ops.outbound_model(
130 self._config, channel, request.model
131 )
132 context = RelayConversionContext(
133 request_id=request.request_id,
134 channel_name=channel.name,
135 upstream_model=outbound_model,
136 options=RelayOptions(),
137 media_resolver=self._media_resolver,
138 )
139 conv = self._converter.convert_request(
140 payload=cast("RelayRequestPayload", request.payload),
141 source=request.source,
142 target=channel.target_format,
143 context=context,
144 )
145 if conv.is_err():
146 if reservation is not None and billing is not None:
147 await billing.release(reservation)
148 return (
149 Err(conversion_error_to_gateway(conv.unwrap_err(), request.request_id)),
150 channel.name,
151 )
152 converted_request = conv.unwrap()
153 session = self._converter.new_stream_session(
154 source=request.source,
155 target=channel.target_format,
156 context=context,
157 )
158 if session.is_err():
159 if reservation is not None and billing is not None:
160 await billing.release(reservation)
161 return (
162 Err(
163 conversion_error_to_gateway(
164 session.unwrap_err(), request.request_id
165 )
166 ),
167 channel.name,
168 )
169 stream_session = session.unwrap()
170 telemetry.log_conversion_loss(
171 request.request_id,
172 converted_request.converter_id,
173 converted_request.losses,
174 )
175 metadata = RelayGatewayMetadata(
176 converter_id=converted_request.converter_id,
177 source=request.source,
178 target=channel.target_format,
179 quality=converted_request.quality,
180 loss_codes=tuple(loss.reason for loss in converted_request.losses),
181 warnings=converted_request.warnings,
182 )
183 stream = self._stream_events(
184 request,
185 channel,
186 outbound_model,
187 converted_request,
188 reservation=reservation,
189 session=stream_session,
190 context=context,
191 )
192 return (
193 Ok(
194 RelayGatewayResult(
195 status_code=200,
196 headers={"x-request-id": request.request_id},
197 payload=None,
198 stream=stream,
199 metadata=metadata,
200 )
201 ),
202 channel.name,
203 )
204
205 async def _stream_events(
206 self,
207 request: RelayGatewayRequest,
208 channel: RelayChannel,
209 outbound_model: str,
210 converted_request: RelayConvertResult[RelayRequestPayload],
211 *,
212 reservation: RelayUsageReservation | None,
213 session: RelayStreamSessionProtocol,
214 context: RelayConversionContext,
215 ) -> AsyncIterator[RelayWireEvent]:
216 """Consume the upstream stream and settle the reservation once.
217
218 Each consumer pull forwards exactly one upstream chunk through
219 the session; cancellation, truncation, and malformed framing
220 follow the ``relay_stream`` lifecycle. The ``finally`` block
221 runs when the consumer ends the stream (completion, disconnect,
222 or error) and settles billing exactly once from the session
223 snapshot.
224
225 Yields:
226 Normalized ``RelayWireEvent`` values framed by the stream
227 session.
228 """
229 streams = self._streams
230 stream_id: str | None = None
231 cancel_handle: asyncio.Event | None = None
232 if streams is not None:
233 stream_id, cancel_handle = streams.register(
234 channel=channel.name,
235 model=outbound_model,
236 request_id=request.request_id,
237 )
238 parser = UpstreamEventParser(
239 session=session,
240 source=channel.target_format,
241 request_id=request.request_id,
242 )
243 url = upstream_ops.upstream_url(channel, outbound_model)
244 payload = (
245 converted_request.value.to_dict()
246 if converted_request.value is not None
247 else {}
248 )
249 logger.info(
250 "relay_gateway_stream_started",
251 request_id=request.request_id,
252 channel=channel.name,
253 method="POST",
254 url=url,
255 )
256 upstream_request = UpstreamRequest(
257 request_id=request.request_id,
258 method="POST",
259 url=url,
260 headers={"content-type": "application/json"},
261 payload=dict(payload),
262 timeout_seconds=channel.timeout_seconds,
263 channel_name=channel.name,
264 )
265 truncated = False
266 stream_iter: AsyncGenerator[RelayWireEvent, None] | None = None
267 try:
268 stream_iter = cast(
269 "AsyncGenerator[RelayWireEvent, None]",
270 relay_stream(
271 cast(
272 "RelayUpstreamProtocol",
273 self._upstream,
274 ),
275 upstream_request,
276 parser,
277 cancel_handle=cancel_handle,
278 ),
279 )
280 try:
281 async for wire in stream_iter:
282 yield wire
283 except (RelayGatewayError, RelayError) as error:
284 logger.warning(
285 "relay_gateway_stream_malformed",
286 request_id=request.request_id,
287 channel=channel.name,
288 error=str(error),
289 )
290 truncated = True
291 raise
292 finally:
293 if stream_iter is not None:
294 await stream_iter.aclose()
295 if truncated:
296 status: Literal["completed", "failed", "cancelled", "truncated"] = (
297 "truncated"
298 )
299 elif parser.cancelled:
300 status = "cancelled"
301 elif parser.truncated:
302 status = "truncated"
303 else:
304 status = "completed"
305 if status == "completed":
306 upstream_ops.note_success(self._failover, channel.name)
307 elif status in ("failed", "truncated"):
308 upstream_ops.note_failure(self._failover, channel.name)
309 if streams is not None and stream_id is not None:
310 streams.unregister(stream_id)
311 billing = self._billing
312 if billing is not None and reservation is not None:
313 settled = billing_ops.stream_settle_result(
314 request,
315 channel,
316 converter_id=converted_request.converter_id,
317 session=session,
318 )
319 await billing_ops.settle(billing, reservation, settled, status=status)
320
321
322__all__ = ["StreamingMixin"]