11"""Generation service — gateway-side adapter over the Inference Engine.
22
3- Week 1 shipped a stub here. Week 2 makes this a thin async adapter that
4- delegates to :class:`inference.engine.InferenceEngine`, preserving the
5- ``generate(request) -> GenerationResponse`` seam the routers depend on.
3+ Week 1 shipped a stub here. Week 2 made this a thin async adapter that delegates
4+ to :class:`inference.engine.InferenceEngine`. Week 4 adds fire-and-forget
5+ inference logging: after a response is produced (or a stream completes) the
6+ service hands an :class:`InferenceRecord` to the :class:`MLflowTracker`, which
7+ schedules the write on a background task. The tracker call never blocks or
8+ raises into the request path, so logging is invisible to latency and to
9+ clients — and a tracking outage cannot fail a generation.
610"""
711
812from __future__ import annotations
913
14+ import time
1015from collections .abc import AsyncIterator
1116
1217from fastapi import Request
1318
1419from inference .engine import InferenceEngine
1520from inference .types import StreamChunk
21+ from shared .context import get_request_id
1622from shared .logging import get_logger
1723from shared .schemas .generation import GenerationRequest , GenerationResponse
24+ from storage .mlflow import InferenceRecord , MLflowTracker
1825
1926logger = get_logger (__name__ )
2027
@@ -24,23 +31,90 @@ class GenerationService:
2431
2532 Args:
2633 engine: The process-wide inference engine (from ``app.state.engine``).
34+ tracker: The process-wide MLflow tracker (from ``app.state.mlflow``).
35+ Defaults to a disabled no-op tracker so the service stays usable
36+ without tracking configured.
2737 """
2838
29- def __init__ (self , engine : InferenceEngine ) -> None :
39+ def __init__ (
40+ self , engine : InferenceEngine , tracker : MLflowTracker | None = None
41+ ) -> None :
3042 self ._engine = engine
43+ self ._tracker = tracker or MLflowTracker .from_settings ()
3144
3245 async def generate (self , request : GenerationRequest ) -> GenerationResponse :
33- """Return a full generation for ``request``."""
34- return await self ._engine .generate (request )
46+ """Return a full generation for ``request`` and log it (best-effort)."""
47+ start = time .perf_counter ()
48+ response = await self ._engine .generate (request )
49+ latency_ms = (time .perf_counter () - start ) * 1000.0
50+
51+ # Fire-and-forget: returns immediately, never raises.
52+ self ._tracker .log_inference (
53+ InferenceRecord (
54+ request_id = get_request_id (),
55+ model = request .model ,
56+ backend = self ._engine .backend_name ,
57+ kind = "generate" ,
58+ max_tokens = request .max_tokens ,
59+ temperature = request .temperature ,
60+ top_p = request .top_p ,
61+ adapter = request .adapter ,
62+ finish_reason = response .finish_reason ,
63+ prompt_tokens = response .usage .prompt_tokens ,
64+ completion_tokens = response .usage .completion_tokens ,
65+ total_tokens = response .usage .total_tokens ,
66+ latency_ms = latency_ms ,
67+ prompt = request .prompt ,
68+ output = response .text ,
69+ )
70+ )
71+ return response
3572
3673 async def generate_stream (
3774 self , request : GenerationRequest
3875 ) -> AsyncIterator [StreamChunk ]:
39- """Yield streaming chunks for ``request``."""
76+ """Yield streaming chunks for ``request``; log once the stream ends.
77+
78+ Accumulates the streamed text and timing (including time-to-first-token)
79+ and emits a single :class:`InferenceRecord` after the terminal chunk.
80+ """
81+ start = time .perf_counter ()
82+ ttft_ms : float | None = None
83+ parts : list [str ] = []
84+ finish_reason : str | None = None
85+
4086 async for chunk in self ._engine .generate_stream (request ):
87+ if chunk .delta :
88+ if ttft_ms is None :
89+ ttft_ms = (time .perf_counter () - start ) * 1000.0
90+ parts .append (chunk .delta )
91+ if chunk .finish_reason is not None :
92+ finish_reason = chunk .finish_reason
4193 yield chunk
4294
95+ latency_ms = (time .perf_counter () - start ) * 1000.0
96+ self ._tracker .log_inference (
97+ InferenceRecord (
98+ request_id = get_request_id (),
99+ model = request .model ,
100+ backend = self ._engine .backend_name ,
101+ kind = "stream" ,
102+ max_tokens = request .max_tokens ,
103+ temperature = request .temperature ,
104+ top_p = request .top_p ,
105+ adapter = request .adapter ,
106+ finish_reason = finish_reason ,
107+ latency_ms = latency_ms ,
108+ ttft_ms = ttft_ms ,
109+ prompt = request .prompt ,
110+ output = "" .join (parts ),
111+ )
112+ )
113+
43114
44115def get_generation_service (request : Request ) -> GenerationService :
45- """FastAPI provider that wires the request's app engine into the service."""
46- return GenerationService (engine = request .app .state .engine )
116+ """FastAPI provider wiring the app's engine + tracker into the service."""
117+ return GenerationService (
118+ engine = request .app .state .engine ,
119+ tracker = getattr (request .app .state , "mlflow" , None ),
120+ )
0 commit comments