Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
11 changes: 8 additions & 3 deletions src/inference/core/handlers/completion.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,7 +14,12 @@
from inference.client import api_gateway_client
from inference.config import settings
from ..providers import resolve_upstream
from ..worker_routing import envoy_route_headers, provider_auth, upstream_model
from ..worker_routing import (
echo_requested_model,
envoy_route_headers,
provider_auth,
upstream_model,
)
from ..rate_limiter import rate_limiter
from ..request_logger import RequestLogger
from ..service import GatewayService
Expand Down Expand Up @@ -194,7 +199,7 @@ def _handle_streaming(
)

processed_stream = StreamProcessor.process_stream(
stream_gen, start_time, tracker
stream_gen, start_time, tracker, rewrite_model=model
)

async def logging_generator_wrapper():
Expand Down Expand Up @@ -288,7 +293,7 @@ async def _handle_standard(
prompt_tokens = usage.get("prompt_tokens", 0)
completion_tokens = usage.get("completion_tokens", 0)

return response_data
return echo_requested_model(response_data, model)
except HTTPException as e:
status_code = e.status_code
error_message = str(e.detail) if hasattr(e, "detail") else str(e)
Expand Down
3 changes: 2 additions & 1 deletion src/inference/core/handlers/embedding.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,7 @@
from fastapi import BackgroundTasks, HTTPException

from ..pipeline import Pipeline, RequestContext
from ..worker_routing import echo_requested_model
from ..request_logger import RequestLogger
from ..service import GatewayService

Expand Down Expand Up @@ -64,7 +65,7 @@ async def handle(
usage = response_data.get("usage", {})
prompt_tokens = usage.get("prompt_tokens", 0)

return response_data
return echo_requested_model(response_data, ctx.model)
except HTTPException as e:
status_code = e.status_code
error_message = str(e.detail) if hasattr(e, "detail") else str(e)
Expand Down
3 changes: 2 additions & 1 deletion src/inference/core/handlers/image.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,7 @@
from fastapi import BackgroundTasks, HTTPException

from ..pipeline import Pipeline, RequestContext
from ..worker_routing import echo_requested_model
from ..request_logger import RequestLogger
from ..service import GatewayService

Expand Down Expand Up @@ -128,7 +129,7 @@ async def _handle(
concurrency_key=ctx.concurrency_key,
)

return response_data
return echo_requested_model(response_data, ctx.model)
except HTTPException as e:
status_code = e.status_code
error_message = str(e.detail) if hasattr(e, "detail") else str(e)
Expand Down
4 changes: 2 additions & 2 deletions src/inference/core/handlers/video.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,7 +11,7 @@
from ..http_client import http_client
from ..pipeline import Pipeline, RequestContext
from ..providers import get_adapter
from ..worker_routing import envoy_route_headers, provider_auth
from ..worker_routing import echo_requested_model, envoy_route_headers, provider_auth
from ..request_logger import RequestLogger
from ..service import GatewayService

Expand Down Expand Up @@ -142,7 +142,7 @@ async def _handle(
timeout=settings.upstream_video_timeout_seconds,
)

return response_data
return echo_requested_model(response_data, ctx.model)
except HTTPException as e:
status_code = e.status_code
error_message = str(e.detail) if hasattr(e, "detail") else str(e)
Expand Down
77 changes: 75 additions & 2 deletions src/inference/core/stream_processor.py
Original file line number Diff line number Diff line change
@@ -1,7 +1,8 @@
import codecs
import json
import logging
import time
from typing import Any, AsyncGenerator, Dict, List
from typing import Any, AsyncGenerator, Dict, List, Optional
from cachetools import LRUCache

logger = logging.getLogger(__name__)
Expand All @@ -17,11 +18,37 @@ class StreamProcessor:
_tiktoken_checked = False
_encoder_cache: LRUCache = LRUCache(maxsize=128)

@staticmethod
def _rewrite_model_line(line: str, model: str) -> str:
"""Replace the `model` field of one SSE line with the requested name.

Anything that is not a JSON `data:` event is returned untouched, which
covers blank separators, comments and the `[DONE]` sentinel.
"""
carriage = "\r" if line.endswith("\r") else ""
stripped = line.rstrip("\r")
if not stripped.startswith("data: "):
return line

payload = stripped[6:].strip()
if not payload or payload == "[DONE]":
return line
try:
event = json.loads(payload)
except json.JSONDecodeError:
return line
if not isinstance(event, dict) or "model" not in event:
return line

event["model"] = model
return f"data: {json.dumps(event, separators=(',', ':'))}{carriage}"

@staticmethod
async def process_stream(
stream_generator: AsyncGenerator,
start_time: float,
usage_tracker: Dict[str, Any],
rewrite_model: Optional[str] = None,
) -> AsyncGenerator[bytes, None]:
"""
Wraps a stream generator to track usage.
Expand All @@ -30,8 +57,38 @@ async def process_stream(
stream_generator: The raw byte stream from upstream
start_time: Request start time (for TTFT)
usage_tracker: Dict to update with 'prompt_tokens', 'completion_tokens', 'ttft_ms'
rewrite_model: When set, every event's `model` is replaced with this
name. Clients address a deployment by name and some OpenAI
client libraries assert the response echoes back what they sent,
but upstream reports its own id. Rewriting forces the stream to
be re-framed into whole lines: upstream is read with
`aiter_raw`, so one `data:` line can arrive split across chunks
and a chunk cannot be rewritten on its own.
"""
buffer = ""
pending = ""
# Incremental, so a multibyte character split across two chunks is held
# until its remaining bytes arrive. A plain bytes.decode(errors="ignore")
# per chunk silently drops the leading part and eats the character.
decoder = codecs.getincrementaldecoder("utf-8")(errors="ignore")

def _framed(text: str, flush: bool = False) -> str:
"""Whole rewritten lines, holding back any trailing fragment."""
nonlocal pending
pending += text
if flush:
out, pending = pending, ""
if not out:
return ""
return StreamProcessor._rewrite_model_line(out, rewrite_model)
lines = pending.split("\n")
pending = lines.pop()
if not lines:
return ""
return "".join(
StreamProcessor._rewrite_model_line(ln, rewrite_model) + "\n"
for ln in lines
)

try:
async for chunk in stream_generator:
Expand All @@ -42,7 +99,18 @@ async def process_stream(
if has_content and usage_tracker.get("ttft_ms") is None:
usage_tracker["ttft_ms"] = int((time.time() - start_time) * 1000)

yield chunk
if rewrite_model is None:
yield chunk
continue

text = (
decoder.decode(chunk)
if isinstance(chunk, bytes)
else str(chunk)
)
out = _framed(text)
if out:
yield out.encode("utf-8")

# Parse any trailing partial line.
if buffer:
Expand All @@ -51,6 +119,11 @@ async def process_stream(
)
if has_content and usage_tracker.get("ttft_ms") is None:
usage_tracker["ttft_ms"] = int((time.time() - start_time) * 1000)

if rewrite_model is not None:
out = _framed(decoder.decode(b"", final=True), flush=True)
if out:
yield out.encode("utf-8")
except Exception as e:
logger.error(f"Stream processing error: {e}")
raise e
Expand Down
21 changes: 20 additions & 1 deletion src/inference/core/worker_routing.py
Original file line number Diff line number Diff line change
Expand Up @@ -27,7 +27,12 @@

logger = logging.getLogger(__name__)

__all__ = ["provider_auth", "upstream_model", "envoy_route_headers"]
__all__ = [
"provider_auth",
"upstream_model",
"echo_requested_model",
"envoy_route_headers",
]

DEPLOYMENT_ID_HEADER = "X-Inferia-Deployment-Id"
ROUTE_CLUSTER_HEADER = "X-Inferia-Route-Cluster"
Expand Down Expand Up @@ -72,6 +77,20 @@ def provider_auth(
return provider_key, {}


def echo_requested_model(response_data: Any, requested: str) -> Any:
"""Put the name the client sent back on an upstream response.

Clients address a deployment by name and the pipeline swaps in the real
upstream id before forwarding, so upstream answers with its own id. Some
OpenAI client libraries assert the response echoes what they sent.

A no-op for engines whose response carries no model field.
"""
if isinstance(response_data, dict) and "model" in response_data:
response_data["model"] = requested
return response_data


def upstream_model(deployment: Dict[str, Any]) -> Optional[str]:
"""The model id to send upstream.

Expand Down
Loading
Loading