diff --git a/packages/sie_gateway/src/handlers/proxy.rs b/packages/sie_gateway/src/handlers/proxy.rs index fb21aa430..e94c3916b 100644 --- a/packages/sie_gateway/src/handlers/proxy.rs +++ b/packages/sie_gateway/src/handlers/proxy.rs @@ -17335,6 +17335,7 @@ mod tests { text: "ok".to_string(), finish_reason: "stop".to_string(), usage: Some(crate::queue::streaming::UsageBlock { + upstream_usage: None, gpu_second: None, images: None, prompt_tokens_details: None, @@ -21196,6 +21197,7 @@ mod tests { text: "Hello world!".to_string(), finish_reason: "stop".to_string(), usage: Some(crate::queue::streaming::UsageBlock { + upstream_usage: None, gpu_second: None, images: None, prompt_tokens_details: None, @@ -26157,6 +26159,7 @@ mod tests { text: String::new(), finish_reason: "stop".to_string(), usage: Some(UsageBlock { + upstream_usage: None, gpu_second: None, images: None, prompt_tokens_details: None, @@ -26215,6 +26218,7 @@ mod tests { text: String::new(), finish_reason: "tool_calls".to_string(), usage: Some(UsageBlock { + upstream_usage: None, gpu_second: None, images: None, prompt_tokens_details: None, @@ -26285,6 +26289,7 @@ mod tests { text: String::new(), finish_reason: "stop".to_string(), usage: Some(UsageBlock { + upstream_usage: None, gpu_second: None, images: None, prompt_tokens_details: None, @@ -26670,6 +26675,7 @@ mod tests { text: "a continuation".to_string(), finish_reason: "length".to_string(), usage: Some(crate::queue::streaming::UsageBlock { + upstream_usage: None, gpu_second: None, images: None, prompt_tokens_details: None, @@ -27296,6 +27302,7 @@ mod tests { text: "a joke".to_string(), finish_reason: "stop".to_string(), usage: Some(crate::queue::streaming::UsageBlock { + upstream_usage: None, gpu_second: None, images: None, prompt_tokens_details: None, @@ -27331,6 +27338,7 @@ mod tests { #[test] fn test_responses_usage_reports_cached_input_tokens() { let usage = crate::queue::streaming::UsageBlock { + upstream_usage: None, gpu_second: None, images: None, prompt_tokens_details: Some(crate::queue::streaming::PromptTokensDetails { @@ -27380,6 +27388,7 @@ mod tests { text: "Hi there!".to_string(), finish_reason: "stop".to_string(), usage: Some(crate::queue::streaming::UsageBlock { + upstream_usage: None, gpu_second: None, images: None, prompt_tokens_details: None, @@ -27473,6 +27482,7 @@ mod tests { text: "Hi".to_string(), finish_reason: "stop".to_string(), usage: Some(crate::queue::streaming::UsageBlock { + upstream_usage: None, gpu_second: None, images: None, prompt_tokens_details: None, @@ -27513,6 +27523,7 @@ mod tests { text: "Hi".to_string(), finish_reason: "stop".to_string(), usage: Some(crate::queue::streaming::UsageBlock { + upstream_usage: None, gpu_second: None, images: None, prompt_tokens_details: None, @@ -27548,6 +27559,7 @@ mod tests { text: String::new(), finish_reason: "tool_calls".to_string(), usage: Some(crate::queue::streaming::UsageBlock { + upstream_usage: None, gpu_second: None, images: None, prompt_tokens_details: None, @@ -28205,6 +28217,7 @@ mod tests { text: "Hi".to_string(), finish_reason: "stop".to_string(), usage: Some(crate::queue::streaming::UsageBlock { + upstream_usage: None, gpu_second: None, images: None, prompt_tokens_details: None, diff --git a/packages/sie_gateway/src/handlers/sse.rs b/packages/sie_gateway/src/handlers/sse.rs index f8faceb59..34ae121fc 100644 --- a/packages/sie_gateway/src/handlers/sse.rs +++ b/packages/sie_gateway/src/handlers/sse.rs @@ -1994,6 +1994,7 @@ mod tests { // whose usage is the count-so-far. finish_reason: "cancelled".to_string(), usage: Some(UsageBlock { + upstream_usage: None, gpu_second: None, images: None, prompt_tokens_details: None, @@ -2304,6 +2305,7 @@ mod tests { let mut terminal = _terminal_chunk("error", None); terminal.seq = 42; terminal.usage = Some(UsageBlock { + upstream_usage: None, gpu_second: None, images: None, prompt_tokens_details: None, @@ -2330,6 +2332,7 @@ mod tests { let mut terminal = _terminal_chunk("error", None); terminal.seq = 42; terminal.usage = Some(UsageBlock { + upstream_usage: None, gpu_second: None, images: None, prompt_tokens_details: None, @@ -3058,6 +3061,7 @@ mod tests { let chunk = _terminal_chunk( "stop", Some(UsageBlock { + upstream_usage: None, gpu_second: None, images: None, prompt_tokens_details: None, @@ -3478,6 +3482,7 @@ mod tests { collector.apply(_terminal_chunk( "stop", Some(UsageBlock { + upstream_usage: None, gpu_second: None, images: None, prompt_tokens_details: None, @@ -3516,6 +3521,7 @@ mod tests { let terminal = _terminal_chunk( "stop", Some(UsageBlock { + upstream_usage: None, gpu_second: None, images: None, prompt_tokens_details: None, @@ -3639,6 +3645,7 @@ mod tests { let terminal = _terminal_chunk( "stop", Some(UsageBlock { + upstream_usage: None, gpu_second: None, images: None, prompt_tokens_details: None, @@ -3687,6 +3694,7 @@ mod tests { let terminal = _terminal_chunk( "stop", Some(UsageBlock { + upstream_usage: None, gpu_second: None, images: None, prompt_tokens_details: None, @@ -3755,6 +3763,7 @@ mod tests { let terminal = _terminal_chunk( "stop", Some(UsageBlock { + upstream_usage: None, gpu_second: None, images: None, prompt_tokens_details: None, @@ -3792,6 +3801,7 @@ mod tests { let terminal = _terminal_chunk( "stop", Some(UsageBlock { + upstream_usage: None, gpu_second: None, images: None, prompt_tokens_details: None, @@ -3927,6 +3937,7 @@ mod tests { collector.apply(_terminal_chunk( "stop", Some(UsageBlock { + upstream_usage: None, gpu_second: None, images: None, prompt_tokens_details: None, diff --git a/packages/sie_gateway/src/queue/streaming.rs b/packages/sie_gateway/src/queue/streaming.rs index 959dac8df..68043fd70 100644 --- a/packages/sie_gateway/src/queue/streaming.rs +++ b/packages/sie_gateway/src/queue/streaming.rs @@ -211,6 +211,26 @@ pub struct UsageBlock { /// engine does not report prefix-cache hits (and from older workers). #[serde(default, skip_serializing_if = "Option::is_none")] pub prompt_tokens_details: Option, + /// The upstream's own counts for a generation it served, when the counts + /// above are the worker's count with the model's tokenizer. Never + /// serialized, so no response surface shows it. + /// + /// `dead_code`-allowed because the `sie-gateway` binary compiles this + /// module tree independently of the library (see the note in `lib.rs`), + /// and only metering consumers of the stream outcome read it. + #[allow(dead_code)] + #[serde(default, skip_serializing)] + pub upstream_usage: Option, +} + +/// Token counts an upstream reported for a generation it served. +#[derive(Debug, Clone, PartialEq, Eq, Deserialize)] +pub struct UpstreamTokenUsage { + pub prompt_tokens: u32, + pub completion_tokens: u32, + /// Clamped to `prompt_tokens` when decoded. + #[serde(default)] + pub cached_tokens: Option, } impl UsageBlock { @@ -240,6 +260,8 @@ struct WorkerUsageBlock { gpu_second: Option, #[serde(default)] prompt_tokens_details: Option, + #[serde(default)] + upstream_usage: Option, } impl From for UsageBlock { @@ -249,6 +271,12 @@ impl From for UsageBlock { .map(|details| PromptTokensDetails { cached_tokens: details.cached_tokens.min(raw.prompt_tokens), }); + let upstream_usage = raw.upstream_usage.map(|upstream| UpstreamTokenUsage { + cached_tokens: upstream + .cached_tokens + .map(|cached| cached.min(upstream.prompt_tokens)), + ..upstream + }); Self { prompt_tokens: raw.prompt_tokens, completion_tokens: raw.completion_tokens, @@ -256,6 +284,7 @@ impl From for UsageBlock { images: raw.images, gpu_second: raw.gpu_second, prompt_tokens_details, + upstream_usage, } } } @@ -1263,6 +1292,7 @@ mod tests { finish_reason: if done { Some("stop".to_string()) } else { None }, usage: if done { Some(UsageBlock { + upstream_usage: None, gpu_second: None, images: None, prompt_tokens_details: None, @@ -1400,6 +1430,64 @@ mod tests { } } + #[test] + fn test_upstream_usage_is_decoded_clamped_and_never_serialized() { + let decode = |upstream_usage: Option| { + let mut usage = serde_json::json!({ + "prompt_tokens": 2, + "completion_tokens": 5, + "total_tokens": 7, + }); + if let Some(value) = upstream_usage { + usage["upstream_usage"] = value; + } + let bytes = rmp_serde::to_vec_named(&usage).expect("encode usage block"); + rmp_serde::from_slice::(&bytes) + }; + + let counted = decode(Some(serde_json::json!({ + "prompt_tokens": 37, + "completion_tokens": 9, + "cached_tokens": 40, + }))) + .expect("upstream usage decodes"); + assert_eq!( + counted.upstream_usage, + Some(UpstreamTokenUsage { + prompt_tokens: 37, + completion_tokens: 9, + cached_tokens: Some(37), + }) + ); + assert_eq!(counted.prompt_tokens, 2); + let serialized = serde_json::to_value(&counted).expect("serialize usage"); + assert!(serialized.get("upstream_usage").is_none()); + assert!(!serialized.to_string().contains("37")); + + let uncached = decode(Some(serde_json::json!({ + "prompt_tokens": 11, + "completion_tokens": 4, + }))) + .expect("upstream usage without a cache count decodes"); + assert_eq!( + uncached + .upstream_usage + .and_then(|usage| usage.cached_tokens), + None + ); + + let legacy = decode(None).expect("usage without an upstream figure decodes"); + assert_eq!(legacy.upstream_usage, None); + + for invalid in [ + serde_json::json!({"prompt_tokens": -1, "completion_tokens": 1}), + serde_json::json!({"prompt_tokens": 1}), + serde_json::json!("37"), + ] { + assert!(decode(Some(invalid)).is_err()); + } + } + #[test] fn test_worker_error_public_contract_fails_closed_for_parameters() { let public_message = "PUBLIC_MESSAGE"; @@ -1750,6 +1838,7 @@ mod tests { collector.last_output_at = Some(first + std::time::Duration::from_millis(400)); collector.output_event_count = 2; collector.final_meta.as_mut().expect("terminal").usage = Some(UsageBlock { + upstream_usage: None, gpu_second: None, images: None, prompt_tokens_details: None, diff --git a/packages/sie_server/REMOTE_BACKENDS.md b/packages/sie_server/REMOTE_BACKENDS.md index d9a5dc36b..4f011aaf8 100644 --- a/packages/sie_server/REMOTE_BACKENDS.md +++ b/packages/sie_server/REMOTE_BACKENDS.md @@ -136,6 +136,20 @@ upstream's `usage` without `input_tokens`. Generation fails closed: a generation, chat or completion response without exact final usage is an error, never a success. +Queued generation for a model with local weights, served by an `openai` +upstream, is counted as if it had run locally. The worker renders the model's +chat template and counts the prompt with the model's tokenizer. It counts the +completion from every text the upstream returned, private reasoning included, +with the same tokenizer. The terminal chunk reports that count without cached +tokens, and carries the upstream's own counts in `usage.upstream_usage`. The +gateway decodes that field for consumers of its stream outcome and never +returns it to a caller. A request whose messages the local template cannot +render, or whose counted prompt exceeds the context length, fails before it is +sent. When the local tokenizer does not load, or has no chat template for a +chat request, the upstream's counts are reported instead. An `sie` upstream +serves the same model and counts with its tokenizer, so its counts are reported +as they are. Single-node serving reports the upstream's counts. + ## OpenAI-compatible embeddings and rerank Define a `kind: openai` upstream with a base URL that includes the provider's diff --git a/packages/sie_server/src/sie_server/adapters/_generation_base.py b/packages/sie_server/src/sie_server/adapters/_generation_base.py index 559287d5b..0f597166b 100644 --- a/packages/sie_server/src/sie_server/adapters/_generation_base.py +++ b/packages/sie_server/src/sie_server/adapters/_generation_base.py @@ -411,6 +411,15 @@ class ToolCallDelta: arguments_delta: str = "" +@dataclass(frozen=True, slots=True) +class UpstreamTokenUsage: + """The token counts an upstream reported for a generation it served.""" + + prompt_tokens: int + completion_tokens: int + cached_tokens: int | None = None + + @dataclass(frozen=True, slots=True) class GenerationChunk: """One chunk yielded by a streaming :meth:`GenerationAdapter.generate`. @@ -472,6 +481,13 @@ class GenerationChunk: # worker forwards it on the wire chunk; the gateway maps it to # ``choices[0].index``. choice_index: int = 0 + # Terminal chunk only: the upstream's own counts when ``prompt_tokens`` + # and ``completion_tokens`` are this server's count of an upstream + # generation. Never published to the caller. + upstream_usage: UpstreamTokenUsage | None = None + # Private reasoning text an upstream returned beside the answer. It is + # counted and then dropped, never published. + reasoning_delta: str = "" # Backwards-compatibility alias: walking-skeleton callers (the local-dev @@ -931,25 +947,39 @@ def unload(self) -> None: # -- Contract ------------------------------------------------------------ async def chat_completion( - self, body: dict[str, Any], *, requested_model: str, max_response_bytes: int = 32 << 20 + self, + body: dict[str, Any], + *, + requested_model: str, + max_response_bytes: int = 32 << 20, + keep_reasoning: bool = False, ) -> dict[str, Any]: """Return a normalized chat answer when this adapter owns chat rendering. The ingress validates and bounds ``body`` before dispatch. Remote adapters pin the upstream model independently from ``requested_model``. ``max_response_bytes`` bounds raw upstream bytes, including discarded - metadata. Local adapters continue through their existing rendering path. + metadata. ``keep_reasoning`` keeps the upstream's private reasoning + text in each message as ``reasoning_content``; callers that set it must + not publish that field. Local adapters continue through their existing + rendering path. """ raise GenerationUnsupportedFieldError("messages", "this generation adapter does not accept chat messages") def chat_completion_stream( - self, body: dict[str, Any], *, requested_model: str, max_response_bytes: int = 32 << 20 + self, + body: dict[str, Any], + *, + requested_model: str, + max_response_bytes: int = 32 << 20, + keep_reasoning: bool = False, ) -> AsyncIterator[dict[str, Any]]: """Stream normalized chat events, including exact final usage. Closing the iterator cancels upstream work. Clean exhaustion certifies every choice, final usage and the upstream's terminal event. The byte bound applies to the complete raw stream, including discarded metadata. + ``keep_reasoning`` behaves as in :meth:`chat_completion`. """ raise GenerationUnsupportedFieldError("messages", "this generation adapter does not accept chat messages") diff --git a/packages/sie_server/src/sie_server/adapters/remote/_chat_transport.py b/packages/sie_server/src/sie_server/adapters/remote/_chat_transport.py index 09ba48a2e..70432af5c 100644 --- a/packages/sie_server/src/sie_server/adapters/remote/_chat_transport.py +++ b/packages/sie_server/src/sie_server/adapters/remote/_chat_transport.py @@ -24,9 +24,10 @@ async def chat_completion( choices: int, max_response_bytes: int, error_body_timeout_s: float, + keep_reasoning: bool = False, ) -> dict[str, Any]: """Return one bounded, normalized chat answer with exact upstream usage.""" - parser = ChatStreamParser(requested_model, choices=choices) + parser = ChatStreamParser(requested_model, choices=choices, keep_reasoning=keep_reasoning) async with open_stream(client, request, upstream=upstream, error_body_timeout_s=error_body_timeout_s) as response: if response.headers.get("content-type", "").partition(";")[0].strip().lower() != "application/json": raise RemoteUpstreamError("upstream did not return a chat answer") @@ -47,9 +48,10 @@ async def chat_completion_stream( choices: int, max_response_bytes: int, error_body_timeout_s: float, + keep_reasoning: bool = False, ) -> AsyncIterator[dict[str, Any]]: """Yield normalized events; upstream failures after output are final.""" - parser = ChatStreamParser(requested_model, choices=choices) + parser = ChatStreamParser(requested_model, choices=choices, keep_reasoning=keep_reasoning) yielded = False try: async with open_stream( diff --git a/packages/sie_server/src/sie_server/adapters/remote/_openai_chat.py b/packages/sie_server/src/sie_server/adapters/remote/_openai_chat.py index e948f4edf..0b694d625 100644 --- a/packages/sie_server/src/sie_server/adapters/remote/_openai_chat.py +++ b/packages/sie_server/src/sie_server/adapters/remote/_openai_chat.py @@ -156,7 +156,7 @@ def _logprobs(value: Any) -> dict[str, Any] | None: return clean -def _message(value: Any, *, stream: bool) -> dict[str, Any]: +def _message(value: Any, *, stream: bool, keep_reasoning: bool = False) -> dict[str, Any]: if not isinstance(value, dict): raise _invalid() clean: dict[str, Any] = {} @@ -170,6 +170,12 @@ def _message(value: Any, *, stream: bool) -> dict[str, Any]: clean[field] = None if text is None else _text(text) if value.get("tool_calls") is not None: clean["tool_calls"] = _tool_calls(value["tool_calls"], stream=stream) + if keep_reasoning: + reasoning = next( + (value[field] for field in ("reasoning_content", "reasoning") if value.get(field) not in (None, "")), None + ) + if reasoning is not None: + clean["reasoning_content"] = _text(reasoning) if not stream and not any(field in clean for field in ("content", "refusal", "tool_calls")): raise _invalid() return clean @@ -178,11 +184,12 @@ def _message(value: Any, *, stream: bool) -> dict[str, Any]: class ChatStreamParser: """Normalize chat events; require each choice to finish and exact final usage.""" - def __init__(self, model: str, *, choices: int = 1) -> None: + def __init__(self, model: str, *, choices: int = 1, keep_reasoning: bool = False) -> None: if isinstance(choices, bool) or not isinstance(choices, int) or not 1 <= choices <= _MAX_CHOICES: raise ValueError("choices must be between 1 and 128") self.model = model self.choices = choices + self.keep_reasoning = keep_reasoning self.done = False self._id = f"chatcmpl-{uuid.uuid4().hex}" self._created = int(time.time()) @@ -256,7 +263,7 @@ def _parse_choices(self, value: Any, *, stream: bool) -> list[dict[str, Any]]: raise _invalid() field = "delta" if stream else "message" raw_message = choice.get(field) - message = _message(raw_message, stream=stream) + message = _message(raw_message, stream=stream, keep_reasoning=self.keep_reasoning) tools = message.get("tool_calls", []) if stream: self._track_tools(index, tools, finished=reason is not None, require_tools=reason == "tool_calls") diff --git a/packages/sie_server/src/sie_server/adapters/remote/openai.py b/packages/sie_server/src/sie_server/adapters/remote/openai.py index 0879aa9d2..a4381b00d 100644 --- a/packages/sie_server/src/sie_server/adapters/remote/openai.py +++ b/packages/sie_server/src/sie_server/adapters/remote/openai.py @@ -225,7 +225,12 @@ def _generation_request( ) async def chat_completion( - self, body: dict[str, Any], *, requested_model: str, max_response_bytes: int = 32 << 20 + self, + body: dict[str, Any], + *, + requested_model: str, + max_response_bytes: int = 32 << 20, + keep_reasoning: bool = False, ) -> dict[str, Any]: client, request = self._generation_request(body, chat=True, stream=False) return await chat_completion( @@ -236,10 +241,16 @@ async def chat_completion( choices=1 if body.get("n") is None else body["n"], max_response_bytes=max_response_bytes, error_body_timeout_s=REQUEST_DEADLINE_S, + keep_reasoning=keep_reasoning, ) def chat_completion_stream( - self, body: dict[str, Any], *, requested_model: str, max_response_bytes: int = 32 << 20 + self, + body: dict[str, Any], + *, + requested_model: str, + max_response_bytes: int = 32 << 20, + keep_reasoning: bool = False, ) -> AsyncIterator[dict[str, Any]]: client, request = self._generation_request(body, chat=True, stream=True) return chat_completion_stream( @@ -250,6 +261,7 @@ def chat_completion_stream( choices=1 if body.get("n") is None else body["n"], max_response_bytes=max_response_bytes, error_body_timeout_s=REQUEST_DEADLINE_S, + keep_reasoning=keep_reasoning, ) def preflight_generate(self, parameters: Mapping[str, Any], *, stream: bool) -> GenerationPreflightResult | None: diff --git a/packages/sie_server/src/sie_server/adapters/remote/sie.py b/packages/sie_server/src/sie_server/adapters/remote/sie.py index 9b66687d7..8fd1a77d7 100644 --- a/packages/sie_server/src/sie_server/adapters/remote/sie.py +++ b/packages/sie_server/src/sie_server/adapters/remote/sie.py @@ -358,7 +358,12 @@ def _chat_request(self, body: dict[str, Any], *, stream: bool) -> tuple[httpx.As return client, request async def chat_completion( - self, body: dict[str, Any], *, requested_model: str, max_response_bytes: int = 32 << 20 + self, + body: dict[str, Any], + *, + requested_model: str, + max_response_bytes: int = 32 << 20, + keep_reasoning: bool = False, ) -> dict[str, Any]: client, request = self._chat_request(body, stream=False) return await chat_completion( @@ -369,10 +374,16 @@ async def chat_completion( choices=1 if body.get("n") is None else body["n"], max_response_bytes=max_response_bytes, error_body_timeout_s=REQUEST_DEADLINE_S, + keep_reasoning=keep_reasoning, ) def chat_completion_stream( - self, body: dict[str, Any], *, requested_model: str, max_response_bytes: int = 32 << 20 + self, + body: dict[str, Any], + *, + requested_model: str, + max_response_bytes: int = 32 << 20, + keep_reasoning: bool = False, ) -> AsyncIterator[dict[str, Any]]: client, request = self._chat_request(body, stream=True) return chat_completion_stream( @@ -383,6 +394,7 @@ def chat_completion_stream( choices=1 if body.get("n") is None else body["n"], max_response_bytes=max_response_bytes, error_body_timeout_s=REQUEST_DEADLINE_S, + keep_reasoning=keep_reasoning, ) def encode( diff --git a/packages/sie_server/src/sie_server/processors/hybrid_usage.py b/packages/sie_server/src/sie_server/processors/hybrid_usage.py new file mode 100644 index 000000000..89c913ba8 --- /dev/null +++ b/packages/sie_server/src/sie_server/processors/hybrid_usage.py @@ -0,0 +1,98 @@ +"""Count an upstream's generation for a model that also has local weights. + +Such a model is counted as if it had run locally. The worker renders the +model's chat template and counts the prompt with the model's tokenizer. This +wrapper counts every generated text the upstream returned, private reasoning +included, with the same tokenizer. The terminal chunk then reports that count, +and the upstream's own counts ride beside it in ``upstream_usage``. +""" + +from __future__ import annotations + +from collections.abc import AsyncIterator, Awaitable, Callable +from dataclasses import dataclass, replace +from typing import Any + +from sie_server.adapters._generation_base import GenerationChunk, UpstreamTokenUsage, aclose_with_error_precedence + + +@dataclass(frozen=True, slots=True) +class HybridCount: + """The worker's own count of one request that an upstream serves. + + ``count_tokens`` returns how many tokens the model's tokenizer makes of a + text. ``completion_limit`` is the most completion tokens the request may + report, over all of its choices. + """ + + prompt_tokens: int + completion_limit: int + count_tokens: Callable[[str], Awaitable[int]] + + +def _only_reasoning(chunk: GenerationChunk) -> bool: + return ( + not chunk.text_delta + and not chunk.done + and chunk.finish_reason is None + and chunk.tool_call_delta is None + and chunk.logprobs is None + and chunk.candidates is None + and chunk.error_code is None + ) + + +def _candidate_texts(candidate: dict[str, Any]) -> list[str]: + texts = [candidate.get("text") or ""] + for call in candidate.get("tool_calls") or (): + function = call.get("function") or {} + texts.extend((function.get("name") or "", function.get("arguments") or "")) + return texts + + +async def count_hybrid_usage( + chunks: AsyncIterator[GenerationChunk], count: HybridCount +) -> AsyncIterator[GenerationChunk]: + """Report the worker's count on the terminal chunk and drop private reasoning. + + A terminal that carries the upstream's counts is reported with + ``count.prompt_tokens``, the counted completion, no cached tokens, and the + upstream's counts in ``upstream_usage``. A completion the upstream counted + but that returned no text counts as one token, and no completion counts + above ``count.completion_limit``. Any other terminal passes through. + """ + texts: dict[int, list[str]] = {} + outcome_selected = False + try: + async for chunk in chunks: + choice = texts.setdefault(chunk.choice_index, []) + choice.extend((chunk.reasoning_delta, chunk.text_delta)) + if chunk.tool_call_delta is not None: + choice.extend((chunk.tool_call_delta.function_name or "", chunk.tool_call_delta.arguments_delta)) + for index, candidate in enumerate(chunk.candidates or ()): + texts.setdefault(index, []).extend(_candidate_texts(candidate)) + if chunk.reasoning_delta: + if _only_reasoning(chunk): + continue + chunk = replace(chunk, reasoning_delta="") + if chunk.done and chunk.prompt_tokens is not None and chunk.completion_tokens is not None: + completion = 0 + for parts in texts.values(): + completion += await count.count_tokens("".join(parts)) + if chunk.completion_tokens > 0: + completion = max(completion, 1) + chunk = replace( + chunk, + prompt_tokens=count.prompt_tokens, + completion_tokens=min(completion, count.completion_limit), + cached_tokens=None, + upstream_usage=UpstreamTokenUsage( + prompt_tokens=chunk.prompt_tokens, + completion_tokens=chunk.completion_tokens, + cached_tokens=chunk.cached_tokens, + ), + ) + outcome_selected = outcome_selected or chunk.done + yield chunk + finally: + await aclose_with_error_precedence(chunks, outcome_selected=outcome_selected, context="hybrid usage count") diff --git a/packages/sie_server/src/sie_server/processors/remote_chat.py b/packages/sie_server/src/sie_server/processors/remote_chat.py index 94955d746..cdd43036d 100644 --- a/packages/sie_server/src/sie_server/processors/remote_chat.py +++ b/packages/sie_server/src/sie_server/processors/remote_chat.py @@ -77,8 +77,13 @@ def remote_chat_chunks( *, requested_model: str, context_length: int, + keep_reasoning: bool = False, ) -> AsyncIterator[GenerationChunk]: - """Prepare chat without dispatch; the worker owns iteration and cancellation.""" + """Prepare chat without dispatch; the worker owns iteration and cancellation. + + With ``keep_reasoning`` the upstream's private reasoning text rides on + ``GenerationChunk.reasoning_delta`` chunks, for the caller to count and drop. + """ if body["stream"]: # Constructing the adapter iterator checks the declared endpoint before # queue admission, but does not send the upstream request. @@ -86,9 +91,10 @@ def remote_chat_chunks( body, requested_model=requested_model, max_response_bytes=_MAX_RESPONSE_BYTES, + keep_reasoning=keep_reasoning, ) return _stream_chunks(iterator, body, context_length) - return _buffered_chunks(adapter, body, requested_model, context_length) + return _buffered_chunks(adapter, body, requested_model, context_length, keep_reasoning) async def _buffered_chunks( @@ -96,19 +102,28 @@ async def _buffered_chunks( body: dict[str, Any], requested_model: str, context_length: int, + keep_reasoning: bool, ) -> AsyncIterator[GenerationChunk]: try: payload = await adapter.chat_completion( body, requested_model=requested_model, max_response_bytes=_MAX_RESPONSE_BYTES, + keep_reasoning=keep_reasoning, ) usage = _usage(payload, body, context_length) if usage is None: raise RemoteUpstreamError("upstream chat omitted exact usage") candidates: list[dict[str, Any]] = [] + reasoning: list[GenerationChunk] = [] for choice in sorted(payload["choices"], key=lambda choice: choice["index"]): message = choice["message"] + if message.get("reasoning_content"): + reasoning.append( + GenerationChunk( + text_delta="", choice_index=choice["index"], reasoning_delta=message["reasoning_content"] + ) + ) tools = message.get("tool_calls") or [] _check_tools({index: tool["function"]["name"] for index, tool in enumerate(tools)}, body, finished=True) if choice["finish_reason"] == "content_filter" or message.get("refusal"): @@ -121,6 +136,8 @@ async def _buffered_chunks( "tool_calls": tools or None, } ) + for chunk in reasoning: + yield chunk if body["n"] > 1: yield _terminal(usage, [choice["finish_reason"] for choice in candidates], candidates=tuple(candidates)) return @@ -166,6 +183,8 @@ async def _stream_chunks( reason = choice["finish_reason"] if reason == "content_filter" or delta.get("refusal"): raise RemoteUpstreamError("upstream chat refused its output") + if delta.get("reasoning_content"): + yield GenerationChunk(text_delta="", choice_index=index, reasoning_delta=delta["reasoning_content"]) calls = tools.setdefault(index, {}) for tool in delta.get("tool_calls") or []: name = tool.get("function", {}).get("name") diff --git a/packages/sie_server/src/sie_server/processors/streaming.py b/packages/sie_server/src/sie_server/processors/streaming.py index df8037711..ad6513224 100644 --- a/packages/sie_server/src/sie_server/processors/streaming.py +++ b/packages/sie_server/src/sie_server/processors/streaming.py @@ -61,6 +61,7 @@ GenerationUnsupportedFieldError, ReasoningFormat, ToolCallDelta, + UpstreamTokenUsage, aclose_with_error_precedence, client_safe_generation_error_code, client_safe_generation_error_message, @@ -88,6 +89,7 @@ from sie_server.processors.generate_params import extract_generate_params from sie_server.processors.grammar_cache import GrammarLRU from sie_server.processors.grammar_compile import compile_outlines +from sie_server.processors.hybrid_usage import HybridCount, count_hybrid_usage from sie_server.processors.remote_chat import remote_chat_chunks from sie_server.processors.strict_grammar import enforce_strict_grammar from sie_server.processors.tool_call_grammar import ( @@ -1888,16 +1890,20 @@ async def _process_inner_guarded( or (isinstance(adapter, OpenAIUpstreamAdapter) and not adapter.supports_raw_completions) ) ) + # A model with local weights that an OpenAI-compatible upstream serves + # is counted with its own tokenizer, as if it had run locally. + counted_upstream = ( + isinstance(adapter, OpenAIUpstreamAdapter) and config is not None and not config.remote_backed + ) + local_template = False if ( isinstance(params.input, _MessagesInput) and isinstance(adapter, (SieUpstreamAdapter, OpenAIUpstreamAdapter)) - and not remote_chat + and (not remote_chat or counted_upstream) ): - try: - tokenizer = await self._get_tokenizer(model_id) - remote_chat = not bool(getattr(tokenizer, "chat_template", None)) - except Exception: # noqa: BLE001 - unavailable local template selects declared upstream chat - remote_chat = True + # An unavailable local template selects declared upstream chat. + local_template = await self._has_chat_template(model_id) + remote_chat = remote_chat or not local_template remote_body: dict[str, Any] | None = None if remote_chat: try: @@ -1924,7 +1930,27 @@ async def _process_inner_guarded( # shape and for text-only message lists. request_images: list[ImageInput] | None = None request_videos: list[VideoInput] | None = None - if remote_body is not None: + if remote_body is not None and counted_upstream and local_template and isinstance(params.input, _MessagesInput): + # Rendered only to count the prompt. The upstream owns chat templating. + rendered_for_count = await self._render_chat_template( + model_id, + params.input.messages, + effective_tools, + request_kwargs=params.chat_template_kwargs, + ) + if isinstance(rendered_for_count, _ValidationError): + await self._terminal_error_then_settle( + reply_subject, + request_id=request_id, + attempt_id=attempt_id, + seq=0, + code=rendered_for_count.code, + message=rendered_for_count.message, + msg=msg, + ) + return + prompt_str = rendered_for_count + elif remote_body is not None: # Used only for the existing admission estimate. The upstream owns # chat templating; remote-backed models have no local tokenizer. prompt_str = json.dumps(remote_body["messages"], ensure_ascii=False) @@ -2013,7 +2039,7 @@ async def _process_inner_guarded( # the context window unaccounted. ctx_error = ( None - if remote_chat + if remote_chat and not (counted_upstream and local_template) else await self._check_context_length( model_id, prompt_str, @@ -2042,6 +2068,31 @@ async def _process_inner_guarded( ) return + hybrid_count: HybridCount | None = None + if counted_upstream and (remote_body is None or local_template): + assert config is not None + assert config.tasks.generate is not None + counted = await self._hybrid_count( + model_id, + prompt_str, + adapter=adapter, + max_new_tokens=params.max_new_tokens, + choices=params.n or 1, + context_length=config.tasks.generate.context_length, + ) + if isinstance(counted, _ValidationError): + await self._terminal_error_then_settle( + reply_subject, + request_id=request_id, + attempt_id=attempt_id, + seq=0, + code=counted.code, + message=counted.message, + msg=msg, + ) + return + hybrid_count = counted + # Resolve the on-wire tool-call format once (config-driven, not a # per-block heuristic) and, for an enforced tool_choice # ("required" / named function), build a constrained-decoding @@ -2128,6 +2179,7 @@ async def _process_inner_guarded( remote_body, requested_model=model_id, context_length=config.tasks.generate.context_length, + keep_reasoning=hybrid_count is not None, ) preflight_result = None else: @@ -2241,6 +2293,7 @@ async def _process_inner_guarded( generation_parameters=generation_parameters, preflight_result=preflight_result, generation_chunks=generation_chunks, + hybrid_count=hybrid_count, cancel_event=cancel_event, ) finally: @@ -2285,6 +2338,7 @@ async def _stream_generate( generation_parameters: Mapping[str, Any], preflight_result: GenerationPreflightResult | None = None, generation_chunks: AsyncIterator[GenerationChunk] | None = None, + hybrid_count: HybridCount | None = None, cancel_event: asyncio.Event, ) -> None: # ``adapter.generate`` is typed as returning ``AsyncIterator`` on @@ -2316,6 +2370,8 @@ async def _stream_generate( if generation_chunks is not None else adapter.generate_with_preflight(gen_kwargs, preflight_result) ) + if hybrid_count is not None: + chunks_iter = count_hybrid_usage(chunks_iter, hybrid_count) if suppress_thinking: chunks_iter = suppress_thinking_blocks( chunks_iter, @@ -2661,6 +2717,7 @@ async def _next_chunk() -> GenerationChunk: prompt_tokens=chunk.prompt_tokens, completion_tokens=chunk.completion_tokens, cached_tokens=chunk.cached_tokens, + upstream_usage=chunk.upstream_usage, images=image_count if not terminal_has_error and chunk.finish_reason != "cancelled" else None, ttft_ms=_compute_ttft_ms(publish_at, first_text_at), error_code=(chunk.error_code or "inference_error") if terminal_has_error else None, @@ -3847,6 +3904,64 @@ async def _check_context_length( ) return None + async def _has_chat_template(self, model_id: str) -> bool: + """Whether the model's tokenizer loads and carries a chat template.""" + try: + tokenizer = await self._get_tokenizer(model_id) + except Exception: # noqa: BLE001 - an unavailable tokenizer has no template + return False + return bool(getattr(tokenizer, "chat_template", None)) + + async def _hybrid_count( + self, + model_id: str, + prompt: str, + *, + adapter: GenerationAdapter, + max_new_tokens: int, + choices: int, + context_length: int, + ) -> HybridCount | _ValidationError | None: + """Count ``prompt`` with the model's tokenizer for an upstream generation. + + The prompt is counted under the same special-token policy as the + context-length guard. ``None`` when the tokenizer does not load: the + generation then reports the upstream's own counts, as a model without + local weights does. A prompt the tokenizer cannot count is refused. + """ + if len(prompt) > _MAX_PROMPT_CHARS: + return _ValidationError( + code="context_exceeded", + message=f"prompt exceeds context_length ({context_length}) for model '{model_id}'", + ) + try: + tok = await self._get_tokenizer(model_id) + except Exception: # noqa: BLE001 - no local tokenizer leaves the upstream's counts + logger.warning("No local tokenizer to count %s; reporting the upstream's counts", model_id, exc_info=True) + return None + loop = asyncio.get_running_loop() + try: + prompt_tokens = await loop.run_in_executor( + _GRAMMAR_EXECUTOR, + lambda: len(tok.encode(prompt, add_special_tokens=adapter.prompt_tokenization_add_special_tokens)), + ) + except Exception: # noqa: BLE001 - an uncountable prompt is refused, never estimated + logger.warning("Prompt count failed for %s", model_id, exc_info=True) + return _ValidationError(code="invalid_request", message=_INTERNAL_CHAT_TEMPLATE_MESSAGE) + + async def count_tokens(text: str) -> int: + if not text: + return 0 + return await loop.run_in_executor( + _GRAMMAR_EXECUTOR, lambda: len(tok.encode(text, add_special_tokens=False)) + ) + + return HybridCount( + prompt_tokens=prompt_tokens, + completion_limit=max(0, min(max_new_tokens, context_length - prompt_tokens)) * choices, + count_tokens=count_tokens, + ) + async def _ensure_grammar_ready( self, grammar: GrammarSpec, @@ -4410,6 +4525,7 @@ def _encode_chunk( prompt_tokens: int | None = None, completion_tokens: int | None = None, cached_tokens: int | None = None, + upstream_usage: UpstreamTokenUsage | None = None, images: int | None = None, ttft_ms: float | None = None, error_code: str | None = None, @@ -4453,6 +4569,14 @@ def _encode_chunk( usage["prompt_tokens_details"] = { "cached_tokens": min(int(cached_tokens), int(prompt_tokens or 0)), } + if upstream_usage is not None: + upstream: dict[str, int] = { + "prompt_tokens": int(upstream_usage.prompt_tokens), + "completion_tokens": int(upstream_usage.completion_tokens), + } + if upstream_usage.cached_tokens is not None: + upstream["cached_tokens"] = min(int(upstream_usage.cached_tokens), int(upstream_usage.prompt_tokens)) + usage["upstream_usage"] = upstream if images is not None: usage["images"] = images payload["usage"] = usage diff --git a/packages/sie_server/src/sie_server/processors/tool_call_parser.py b/packages/sie_server/src/sie_server/processors/tool_call_parser.py index 4f0d786fa..85bd1569f 100644 --- a/packages/sie_server/src/sie_server/processors/tool_call_parser.py +++ b/packages/sie_server/src/sie_server/processors/tool_call_parser.py @@ -649,6 +649,7 @@ def _state(idx: int) -> _ChoiceState: prompt_tokens=chunk.prompt_tokens, completion_tokens=chunk.completion_tokens, cached_tokens=chunk.cached_tokens, + upstream_usage=chunk.upstream_usage, candidates=tuple(updated_candidates), logprobs=chunk.logprobs, error_code=chunk.error_code, @@ -698,6 +699,7 @@ def _state(idx: int) -> _ChoiceState: prompt_tokens=chunk.prompt_tokens, completion_tokens=chunk.completion_tokens, cached_tokens=chunk.cached_tokens, + upstream_usage=chunk.upstream_usage, error_code=chunk.error_code, error_message=chunk.error_message, # Preserve ``candidates`` (with any per-candidate diff --git a/packages/sie_server/tests/adapters/test_remote_openai_chat.py b/packages/sie_server/tests/adapters/test_remote_openai_chat.py index c3b88fec9..4b1147b4e 100644 --- a/packages/sie_server/tests/adapters/test_remote_openai_chat.py +++ b/packages/sie_server/tests/adapters/test_remote_openai_chat.py @@ -293,3 +293,21 @@ def test_discarded_reasoning_cannot_leak_through_logprobs(stream: bool, reasonin assert result is not None assert result["choices"][0]["logprobs"] is None assert "private reasoning" not in json.dumps(result) + + +@pytest.mark.parametrize("field", ["reasoning_content", "reasoning"]) +def test_reasoning_is_kept_only_on_request(field: str) -> None: + kept = ChatStreamParser("model", keep_reasoning=True).completion(wire([choice(**{field: "private"})], usage=USAGE)) + assert kept["choices"][0]["message"] == {"role": "assistant", "content": "answer", "reasoning_content": "private"} + dropped = ChatStreamParser("model").completion(wire([choice(**{field: "private"})], usage=USAGE)) + assert "private" not in json.dumps(dropped) + streamed = ChatStreamParser("model", keep_reasoning=True).parse( + wire([choice(stream=True, finish=None, content=None, **{field: "step"})]) + ) + assert streamed is not None + assert streamed["choices"][0]["delta"]["reasoning_content"] == "step" + + +def test_kept_reasoning_must_be_text() -> None: + with pytest.raises(RemoteUpstreamError): + ChatStreamParser("model", keep_reasoning=True).completion(wire([choice(reasoning_content=7)], usage=USAGE)) diff --git a/packages/sie_server/tests/processors/test_hybrid_usage.py b/packages/sie_server/tests/processors/test_hybrid_usage.py new file mode 100644 index 000000000..69d8ada5b --- /dev/null +++ b/packages/sie_server/tests/processors/test_hybrid_usage.py @@ -0,0 +1,412 @@ +import asyncio +import json +import time +from collections.abc import AsyncIterator +from typing import Any +from unittest.mock import AsyncMock, MagicMock + +import httpx +import msgpack +import pytest +from sie_server.adapters._generation_base import GenerationChunk, ToolCallDelta, UpstreamTokenUsage +from sie_server.adapters.remote.openai import OpenAIUpstreamAdapter +from sie_server.adapters.remote.sie import SieUpstreamAdapter +from sie_server.config.model import ModelConfig +from sie_server.config.upstreams import Upstream, install_upstreams +from sie_server.core.upstream_client import upstream_client +from sie_server.processors.hybrid_usage import HybridCount, count_hybrid_usage +from sie_server.processors.streaming import StreamingProcessor + +MODEL = "local/hybrid" +UPSTREAM_USAGE = { + "prompt_tokens": 37, + "completion_tokens": 9, + "total_tokens": 46, + "prompt_tokens_details": {"cached_tokens": 30}, +} + + +async def _words(text: str) -> int: + return len(text.split()) + + +def _count(*, prompt_tokens: int = 2, completion_limit: int = 100) -> HybridCount: + return HybridCount(prompt_tokens=prompt_tokens, completion_limit=completion_limit, count_tokens=_words) + + +class _Chunks: + def __init__(self, chunks: list[GenerationChunk]) -> None: + self._chunks = iter(chunks) + self.closed = False + + def __aiter__(self) -> "_Chunks": + return self + + async def __anext__(self) -> GenerationChunk: + try: + return next(self._chunks) + except StopIteration: + raise StopAsyncIteration from None + + async def aclose(self) -> None: + self.closed = True + + +async def _drain(chunks: list[GenerationChunk], count: HybridCount) -> tuple[list[GenerationChunk], _Chunks]: + source = _Chunks(chunks) + return [chunk async for chunk in count_hybrid_usage(source, count)], source + + +def _terminal(**fields: Any) -> GenerationChunk: + terminal: dict[str, Any] = { + "text_delta": "", + "done": True, + "finish_reason": "stop", + "prompt_tokens": 37, + "completion_tokens": 9, + "cached_tokens": 30, + } + return GenerationChunk(**{**terminal, **fields}) + + +async def test_terminal_reports_the_worker_count_and_keeps_the_upstream_figure_beside_it() -> None: + out, source = await _drain( + [ + GenerationChunk(text_delta="", reasoning_delta="think one two "), + GenerationChunk(text_delta="answer here", is_first=True), + _terminal(), + ], + _count(), + ) + assert [chunk.text_delta for chunk in out] == ["answer here", ""] + assert all(not chunk.reasoning_delta for chunk in out) + terminal = out[-1] + assert (terminal.prompt_tokens, terminal.completion_tokens, terminal.cached_tokens) == (2, 5, None) + assert terminal.upstream_usage == UpstreamTokenUsage(prompt_tokens=37, completion_tokens=9, cached_tokens=30) + assert source.closed + + +async def test_reasoning_on_a_visible_chunk_is_counted_and_stripped() -> None: + out, _ = await _drain( + [GenerationChunk(text_delta="answer", reasoning_delta="why "), _terminal()], + _count(), + ) + assert out[0] == GenerationChunk(text_delta="answer") + assert out[-1].completion_tokens == 2 + + +async def test_tool_calls_and_candidates_are_counted_per_choice() -> None: + call = ToolCallDelta(index=0, id="call-a", function_name="lookup", arguments_delta='{"q": "x"}') + streamed, _ = await _drain( + [GenerationChunk(text_delta="", tool_call_delta=call), _terminal(finish_reason="tool_calls")], + _count(), + ) + assert streamed[-1].completion_tokens == 2 + candidates = ( + {"text": "first answer", "finish_reason": "stop"}, + { + "text": "", + "finish_reason": "tool_calls", + "tool_calls": [{"id": "c", "type": "function", "function": {"name": "lookup", "arguments": "{}"}}], + }, + ) + buffered, _ = await _drain( + [GenerationChunk(text_delta="", choice_index=1, reasoning_delta="private "), _terminal(candidates=candidates)], + _count(), + ) + assert buffered[-1].completion_tokens == 4 + + +async def test_completion_is_bounded_and_never_zero_when_the_upstream_counted_one() -> None: + bounded, _ = await _drain([GenerationChunk(text_delta="a b c d e"), _terminal()], _count(completion_limit=3)) + assert bounded[-1].completion_tokens == 3 + empty, _ = await _drain([_terminal()], _count()) + assert empty[-1].completion_tokens == 1 + + +async def test_a_terminal_without_upstream_counts_passes_through() -> None: + cancelled = GenerationChunk(text_delta="", done=True, finish_reason="cancelled") + out, _ = await _drain([GenerationChunk(text_delta="partial"), cancelled], _count()) + assert out[-1] == cancelled + + +class WordTokenizer: + chat_template = "template" + + def __init__(self, *, fail: bool = False, words: int | None = None) -> None: + self.fail = fail + self.words = words + + def apply_chat_template(self, messages: list[dict], *, tokenize: bool, add_generation_prompt: bool, **kwargs: Any): + assert tokenize is False + assert add_generation_prompt is True + if self.fail: + raise ValueError("template rejects these messages") + if self.words is not None: + return "w " * self.words + rendered = " ".join(message["content"] for message in messages) + return rendered + (" tools" if kwargs.get("tools") else "") + " reply:" + + def encode(self, text: str, add_special_tokens: bool = False) -> list[str]: + return text.split() + + +class ChatStream(httpx.AsyncByteStream): + def __init__(self, frames: list[dict | str]) -> None: + self.frames = frames + self.closed = False + + async def __aiter__(self) -> AsyncIterator[bytes]: + for frame in self.frames: + data = frame if isinstance(frame, str) else json.dumps(frame) + yield f"data: {data}\n\n".encode() + + async def aclose(self) -> None: + self.closed = True + + +def _frames() -> list[dict | str]: + return [ + {"choices": [{"index": 0, "delta": {"reasoning_content": "think one two "}, "finish_reason": None}]}, + {"choices": [{"index": 0, "delta": {"content": "answer here"}, "finish_reason": None}]}, + {"choices": [{"index": 0, "delta": {}, "finish_reason": "stop"}]}, + {"choices": [], "usage": UPSTREAM_USAGE}, + "[DONE]", + ] + + +def _answer() -> dict: + message = {"role": "assistant", "content": "answer here", "reasoning_content": "think one two "} + return {"choices": [{"index": 0, "message": message, "finish_reason": "stop"}], "usage": UPSTREAM_USAGE} + + +def _processor(kind: str, endpoints: list[str]) -> tuple: + cls = SieUpstreamAdapter if kind == "sie" else OpenAIUpstreamAdapter + upstream = Upstream.model_validate( + { + "kind": kind, + "base_url": "http://127.0.0.1:8088/prefix", + **({"endpoints": endpoints} if kind == "openai" else {}), + "rate_cap": {"requests_per_minute": 600, "max_concurrency": 8}, + } + ) + install_upstreams({"hybrid-chat": upstream}) + adapter = cls(upstream="hybrid-chat", upstream_model="operator/model") + adapter.load("cpu") + config = ModelConfig.model_validate( + { + "sie_id": MODEL, + "hf_id": "local/hybrid", + "hf_revision": "a" * 40, + "inputs": {"text": True}, + "tasks": {"generate": {"context_length": 4096, "max_output_tokens": 64, "capabilities": {"tools": True}}}, + "profiles": { + "default": { + "adapter_path": f"{cls.__module__}:{cls.__name__}", + "max_batch_tokens": 8192, + "adapter_options": {"loadtime": {"upstream": "hybrid-chat", "upstream_model": "operator/model"}}, + } + }, + } + ) + registry = MagicMock() + registry.device = "cpu" + registry.is_loaded.return_value = True + registry.get.return_value = adapter + registry.get_config.return_value = config + nc = AsyncMock() + proc = StreamingProcessor(nc=nc, registry=registry, worker_id="w1") + tokenizer = WordTokenizer() + proc._get_tokenizer = AsyncMock(return_value=tokenizer) # type: ignore[method-assign] + return adapter, proc, nc, [], tokenizer + + +async def _close(adapter: Any) -> None: + await adapter.aclose_client() + adapter.unload() + install_upstreams({}) + + +@pytest.fixture +async def hybrid() -> AsyncIterator[tuple]: + setup = _processor("openai", ["chat"]) + try: + yield setup + finally: + await _close(setup[0]) + + +def respond(hybrid: tuple, *, frames: list[dict | str] | None = None, payload: dict | None = None) -> None: + adapter, _proc, _nc, requests, _tokenizer = hybrid + + def handler(request: httpx.Request) -> httpx.Response: + requests.append(request) + if payload is not None: + return httpx.Response( + 200, headers={"content-type": "application/json"}, stream=httpx.ByteStream(json.dumps(payload).encode()) + ) + return httpx.Response( + 200, headers={"content-type": "text/event-stream"}, stream=ChatStream(frames or _frames()) + ) + + adapter._async_client = upstream_client(adapter._upstream, transport=httpx.MockTransport(handler)) + + +async def run(hybrid: tuple, **fields: Any) -> list[dict]: + _adapter, proc, nc, _requests, _tokenizer = hybrid + msg = AsyncMock() + msg.data = msgpack.packb( + { + "request_id": "req-1", + "work_item_id": "req-1.0", + "item_index": 0, + "total_items": 1, + "operation": "generate", + "model_id": MODEL, + "profile_id": "default", + "pool_name": "default", + "router_id": "router-1", + "reply_subject": "_INBOX.router-1.req-1", + "timestamp": time.time(), + "generate": { + "messages": [{"role": "user", "content": "question"}], + "max_new_tokens": 9, + "stream": True, + **fields, + }, + }, + use_bin_type=True, + ) + await proc.process(msg, MODEL) + msg.ack.assert_awaited_once() + return [msgpack.unpackb(call.args[1], raw=False) for call in nc.publish.await_args_list] + + +@pytest.mark.parametrize("streaming", [True, False]) +async def test_queued_upstream_chat_reports_the_worker_count(hybrid: tuple, streaming: bool) -> None: + if streaming: + respond(hybrid) + else: + respond(hybrid, payload=_answer()) + chunks = await run(hybrid, stream=streaming) + assert "".join(chunk.get("text_delta", "") for chunk in chunks) == "answer here" + assert "think" not in json.dumps(chunks) + assert chunks[-1]["finish_reason"] == "stop" + assert chunks[-1]["usage"] == { + "prompt_tokens": 2, + "completion_tokens": 5, + "total_tokens": 7, + "upstream_usage": {"prompt_tokens": 37, "completion_tokens": 9, "cached_tokens": 30}, + } + sent = json.loads(hybrid[3][0].content) + assert sent["messages"] == [{"role": "user", "content": "question"}] + + +async def test_a_platform_upstream_reports_its_own_count() -> None: + setup = _processor("sie", []) + setup[1]._get_tokenizer.side_effect = AssertionError("a platform upstream counts with its own tokenizer") + frames = [ + {"choices": [{"index": 0, "delta": {"content": "first"}, "finish_reason": "stop"}]}, + {"choices": [{"index": 1, "delta": {"content": "second"}, "finish_reason": "stop"}]}, + {"choices": [], "usage": UPSTREAM_USAGE}, + "[DONE]", + ] + try: + respond(setup, frames=frames) + chunks = await run(setup, n=2) + finally: + await _close(setup[0]) + assert chunks[-1]["usage"] == UPSTREAM_USAGE + + +async def test_queued_upstream_chat_renders_tools_into_the_counted_prompt(hybrid: tuple) -> None: + respond(hybrid) + tool = {"type": "function", "function": {"name": "lookup", "parameters": {"type": "object"}}} + chunks = await run(hybrid, tools=[tool]) + assert chunks[-1]["usage"]["prompt_tokens"] == 3 + + +@pytest.mark.parametrize( + ("tokenizer", "code"), + [(WordTokenizer(fail=True), "invalid_request"), (WordTokenizer(words=4090), "context_exceeded")], +) +async def test_a_prompt_the_worker_cannot_count_is_refused_before_dispatch( + hybrid: tuple, monkeypatch: pytest.MonkeyPatch, tokenizer: WordTokenizer, code: str +) -> None: + monkeypatch.setattr(hybrid[1], "_get_tokenizer", AsyncMock(return_value=tokenizer)) + respond(hybrid) + chunks = await run(hybrid) + assert chunks[-1]["error"]["code"] == code + assert hybrid[3] == [] + + +async def test_without_a_local_tokenizer_the_upstream_counts_are_reported(hybrid: tuple) -> None: + hybrid[1]._get_tokenizer.side_effect = RuntimeError("no local tokenizer") + respond(hybrid) + chunks = await run(hybrid) + assert chunks[-1]["usage"] == UPSTREAM_USAGE + assert "think" not in json.dumps(chunks) + + +async def test_a_remote_backed_model_reports_the_upstream_figure(hybrid: tuple) -> None: + config = hybrid[1]._registry.get_config.return_value + data = config.model_dump() + data.update(remote_backed=True, hf_id=None, hf_revision=None) + hybrid[1]._registry.get_config.return_value = ModelConfig.model_validate(data) + hybrid[1]._get_tokenizer.side_effect = AssertionError("a remote-backed model has no local tokenizer") + respond(hybrid) + chunks = await run(hybrid) + assert chunks[-1]["usage"] == UPSTREAM_USAGE + + +@pytest.mark.parametrize("tools", [False, True]) +async def test_upstream_raw_completion_reports_the_worker_count(tools: bool) -> None: + setup = _processor("openai", ["chat", "completions"]) + tool = {"type": "function", "function": {"name": "lookup", "parameters": {"type": "object"}}} + frames = [ + {"choices": [{"index": 0, "text": "one two ", "finish_reason": None}]}, + {"choices": [{"index": 0, "text": "three", "finish_reason": "stop"}]}, + {"choices": [], "usage": {"prompt_tokens": 11, "completion_tokens": 4, "total_tokens": 15}}, + "[DONE]", + ] + try: + respond(setup, frames=frames) + chunks = await run(setup, **({"tools": [tool]} if tools else {})) + finally: + await _close(setup[0]) + requests = setup[3] + prompt = "question tools reply:" if tools else "question reply:" + assert requests[0].url.path == "/prefix/completions" + assert json.loads(requests[0].content)["prompt"] == prompt + assert chunks[-1]["usage"] == { + "prompt_tokens": len(prompt.split()), + "completion_tokens": 3, + "total_tokens": len(prompt.split()) + 3, + "upstream_usage": {"prompt_tokens": 11, "completion_tokens": 4}, + } + + +async def test_cancellation_closes_the_counted_upstream_chat(hybrid: tuple) -> None: + started = asyncio.Event() + + class Held(ChatStream): + async def __aiter__(self) -> AsyncIterator[bytes]: + yield f"data: {json.dumps(_frames()[1])}\n\n".encode() + started.set() + await asyncio.Event().wait() + + stream = Held([]) + adapter, proc = hybrid[0], hybrid[1] + + def handler(request: httpx.Request) -> httpx.Response: + hybrid[3].append(request) + return httpx.Response(200, headers={"content-type": "text/event-stream"}, stream=stream) + + adapter._async_client = upstream_client(adapter._upstream, transport=httpx.MockTransport(handler)) + task = asyncio.create_task(run(hybrid)) + await asyncio.wait_for(started.wait(), timeout=2) + assert proc.signal_cancel("req-1") + chunks = await asyncio.wait_for(task, timeout=2) + assert chunks[-1]["finish_reason"] == "cancelled" + assert "usage" not in chunks[-1] + assert stream.closed diff --git a/packages/sie_server/tests/processors/test_remote_chat.py b/packages/sie_server/tests/processors/test_remote_chat.py index 485e34968..694942d24 100644 --- a/packages/sie_server/tests/processors/test_remote_chat.py +++ b/packages/sie_server/tests/processors/test_remote_chat.py @@ -568,7 +568,7 @@ async def test_onboarded_strict_tools_keep_declared_chat(remote: tuple, monkeypa config_data = proc._registry.get_config(MODEL).model_dump() config_data.update({"remote_backed": False, "hf_id": "local/model"}) proc._registry.get_config.return_value = ModelConfig.model_validate(config_data) - tokenizer = AsyncMock(side_effect=AssertionError("strict tools must retain chat ownership")) + tokenizer = AsyncMock(side_effect=RuntimeError("no local tokenizer")) monkeypatch.setattr(proc, "_get_tokenizer", tokenizer) tool = {**TOOL, "function": {**TOOL["function"], "strict": True}} respond(remote) @@ -576,7 +576,8 @@ async def test_onboarded_strict_tools_keep_declared_chat(remote: tuple, monkeypa assert chunks[-1]["finish_reason"] == "stop" assert requests[0].url.path.endswith("/chat/completions") assert json.loads(requests[0].content)["tools"][0]["function"]["strict"] is True - tokenizer.assert_not_awaited() + counts = isinstance(adapter, OpenAIUpstreamAdapter) + assert tokenizer.await_count == int(counts), "only an OpenAI upstream consults the tokenizer, to count" @pytest.mark.parametrize("remote", ["sie"], indirect=True)