diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml index 9d80fe6ff5..1919fac9c5 100644 --- a/.pre-commit-config.yaml +++ b/.pre-commit-config.yaml @@ -1,6 +1,6 @@ repos: - repo: https://github.com/astral-sh/ruff-pre-commit - rev: v0.15.13 + rev: v0.15.14 hooks: - id: ruff args: diff --git a/studio/backend/core/export/export.py b/studio/backend/core/export/export.py index 4ab95d896f..7cabd382eb 100644 --- a/studio/backend/core/export/export.py +++ b/studio/backend/core/export/export.py @@ -475,6 +475,7 @@ class ExportBackend: self.current_model.save_pretrained_merged( save_directory, self.current_tokenizer, + save_method = "merged_16bit", ) else: self.current_model.save_pretrained(save_directory) @@ -510,6 +511,7 @@ class ExportBackend: self.current_model.save_pretrained_merged( tmp_dir, self.current_tokenizer, + save_method = "merged_16bit", ) self.current_model.push_to_hub_merged( repo_id, diff --git a/studio/backend/core/inference/external_provider.py b/studio/backend/core/inference/external_provider.py index d1f2aec813..1cdf0a19f5 100644 --- a/studio/backend/core/inference/external_provider.py +++ b/studio/backend/core/inference/external_provider.py @@ -112,6 +112,136 @@ _ANTHROPIC_4_7_SAMPLING_REMOVED = re.compile( ) _OPENAI_REASONING_SUMMARY_UNSUPPORTED = re.compile(r"^o3(?:[-.]|$)") +# OpenAI Responses inline citation markers: `citeSOURCE_ID[id2...][LOCATOR]` +# using private-use codepoints (see +# https://developers.openai.com/api/docs/guides/citation-formatting). +# Group 1 holds the delim-separated tokens; each resolvable token expands +# to `[[N]](URL)`, unresolved tokens (locators, unknown ids) drop silently +# so no garbled glyph reaches the renderer. +_OPENAI_CITE_OPEN = "cite" +_OPENAI_CITE_STOP = "" +_OPENAI_CITE_DELIM = "" +_OPENAI_CITATION_MARKER = re.compile( + f"{_OPENAI_CITE_OPEN}([^{_OPENAI_CITE_STOP}]+){_OPENAI_CITE_STOP}" +) + + +def _build_citation_lookup( + url_citations: list[dict[str, Any]], +) -> dict[str, tuple[int, str]]: + """Map every known ``source_id`` alias to ``(citation_index, url)``. + + Accepts singular ``source_id`` and plural ``source_ids``. First-seen + wins on alias collision so an earlier citation keeps its number. + """ + by_source: dict[str, tuple[int, str]] = {} + for idx, cit in enumerate(url_citations, start = 1): + url = cit.get("url") + if not isinstance(url, str) or not url: + continue + aliases: list[str] = [] + sid = cit.get("source_id") + if isinstance(sid, str) and sid: + aliases.append(sid) + sids = cit.get("source_ids") + if isinstance(sids, list): + aliases.extend(s for s in sids if isinstance(s, str) and s) + for alias in aliases: + by_source.setdefault(alias, (idx, url)) + return by_source + + +def _replace_openai_citation_markers( + text: str, + url_citations: list[dict[str, Any]], +) -> str: + """Rewrite `\\ue200cite\\ue202SOURCE_ID[\\ue202LOCATOR]\\ue201` markers into + `[[N]](URL)` per resolvable id. Multi-source markers expand to one link + per id; unresolved tokens drop silently. Idempotent on text without + private-use codepoints. + """ + if not text or _OPENAI_CITE_STOP not in text: + return text + by_source = _build_citation_lookup(url_citations) + + def _sub(match: re.Match[str]) -> str: + # Try every delim-split token; unresolved tokens drop silently. + # Handles multi-source (all resolve) and source+locator (only the + # id resolves, locator drops). Empty result strips the marker. + rendered: list[str] = [] + for tok in match.group(1).split(_OPENAI_CITE_DELIM): + if not tok: + continue + hit = by_source.get(tok) + if hit is None: + continue + idx, url = hit + rendered.append(f"[[{idx}]]({url})") + return "".join(rendered) + + return _OPENAI_CITATION_MARKER.sub(_sub, text) + + +def _rewrite_citation_markers_partial( + text: str, + url_citations: list[dict[str, Any]], +) -> tuple[str, bool]: + """Like ``_replace_openai_citation_markers`` but also reports whether + any marker referenced a source_id not yet in ``url_citations``. + + The ``annotation.added`` event for a url_citation typically arrives + AFTER the delta carrying the marker referencing it. Callers buffer the + segment until a later event records the annotation; unresolved markers + are left verbatim so a follow-up pass still parses cleanly. + """ + if not text or _OPENAI_CITE_STOP not in text: + return text, False + by_source = _build_citation_lookup(url_citations) + has_unresolved = False + + def _sub(match: re.Match[str]) -> str: + nonlocal has_unresolved + tokens = [t for t in match.group(1).split(_OPENAI_CITE_DELIM) if t] + rendered: list[str] = [] + any_unresolved = False + for tok in tokens: + hit = by_source.get(tok) + if hit is None: + any_unresolved = True + continue + idx, url = hit + rendered.append(f"[[{idx}]]({url})") + # Leave the whole marker verbatim if any token is unresolved so the + # caller can re-run once the late annotation lands; partial emission + # would lose the unresolved ids once the source text is dropped. + if any_unresolved: + has_unresolved = True + return match.group(0) + return "".join(rendered) + + return _OPENAI_CITATION_MARKER.sub(_sub, text), has_unresolved + + +def _split_pending_citation_tail(text: str) -> tuple[str, str]: + """Split ``text`` into ``(head, pending_tail)`` for streamed deltas. + + A citation marker can straddle two SSE deltas (e.g. delta-1 ends with + ``\\ue200citetu`` and delta-2 starts with ``rn0view0\\ue201``); the + unterminated tail is buffered and prepended onto the next delta so the + rewriter sees a complete marker. ``pending_tail`` is the longest suffix + starting with ``\\ue200`` and lacking ``\\ue201``; ``head`` is safe to + emit. Empty tail when ``text`` has no open marker or a fully closed one. + """ + if not text: + return text, "" + last_open = text.rfind("") + if last_open == -1: + return text, "" + # Stop byte after the last open byte means the marker closed in this delta. + if _OPENAI_CITE_STOP in text[last_open:]: + return text, "" + return text[:last_open], text[last_open:] + class _AnthropicThinkingSpec(NamedTuple): prefixes: tuple[str, ...] @@ -217,10 +347,87 @@ _ANTHROPIC_COMPACTION_TYPE = "compact_20260112" _ANTHROPIC_COMPACTION_MIN = 50_000 +# Anthropic fast-mode beta (Opus 4.6 / 4.7 only, per +# https://platform.claude.com/docs/en/build-with-claude/fast-mode). +# Mutually exclusive with the Priority service tier. +_ANTHROPIC_FAST_MODE_BETA = "fast-mode-2026-02-01" +_ANTHROPIC_FAST_MODE_PREFIXES = ( + "claude-opus-4-7", + "claude-opus-4-6", +) + + def _anthropic_supports_compaction(model: str) -> bool: return model.startswith(_ANTHROPIC_COMPACTION_PREFIXES) +def _anthropic_supports_fast_mode(model: str) -> bool: + # Require a family boundary ("" or "-") after the prefix so IDs like + # "claude-opus-4-70" / "claude-opus-4-7b" do not match. + return any( + model == p or model.startswith(f"{p}-") for p in _ANTHROPIC_FAST_MODE_PREFIXES + ) + + +# Cap on ``cited_text`` forwarded in document_citations tool_events; +# keeps SSE bytes bounded on multi-KB cited spans (frontend trims to +# 240 chars anyway). +_CITED_TEXT_MAX_LEN = 512 + + +def _anthropic_citation_key(citation: dict[str, Any]) -> tuple: + """Stable dedup key for an Anthropic ``citations_delta.citation``. + + Anchor fields vary per type (char_location, page_location, + content_block_location, search_result_location); both start AND + exclusive end indices are part of the key so same-start / + different-end pairs stay distinct. search_result_location keys on + ``search_result_index`` + ``source`` instead of document_index so + distinct results with the same source don't collapse. Unknown + shapes fall back to a stringified copy (more entries, never + collisions). See + https://platform.claude.com/docs/en/build-with-claude/citations + and https://platform.claude.com/docs/en/build-with-claude/search-results. + """ + ctype = citation.get("type") + doc = citation.get("document_index") + title = citation.get("document_title") or "" + if ctype == "char_location": + return ( + ctype, + doc, + title, + citation.get("start_char_index"), + citation.get("end_char_index"), + ) + if ctype == "page_location": + return ( + ctype, + doc, + title, + citation.get("start_page_number"), + citation.get("end_page_number"), + ) + if ctype == "content_block_location": + return ( + ctype, + doc, + title, + citation.get("start_block_index"), + citation.get("end_block_index"), + ) + if ctype == "search_result_location": + return ( + ctype, + citation.get("search_result_index"), + citation.get("source"), + citation.get("title") or "", + citation.get("start_block_index"), + citation.get("end_block_index"), + ) + return (ctype, _json.dumps(citation, sort_keys = True)) + + class _MistralThinkingSpec(NamedTuple): models: tuple[str, ...] style: Literal["prompt_mode", "reasoning_effort", "disabled"] @@ -397,6 +604,7 @@ class ExternalProviderClient: stop: Optional[Union[str, list[str]]] = None, service_tier: Optional[str] = None, parallel_tool_calls: Optional[bool] = None, + fast_mode: Optional[bool] = None, stream: bool = True, ) -> AsyncGenerator[str, None]: """ @@ -415,6 +623,9 @@ class ExternalProviderClient: stream helpers silently drop fields the upstream API does not accept (e.g. Responses rejects all of seed / frequency / stop; Anthropic does not implement seed / frequency / logprobs). + + ``fast_mode`` only applies to Anthropic Opus 4.6 / 4.7 (silently + dropped elsewhere); adds the beta header and ``speed: "fast"``. """ if not self._is_openai_compatible(): async for line in self._stream_anthropic( @@ -434,6 +645,7 @@ class ExternalProviderClient: stop = stop, service_tier = service_tier, parallel_tool_calls = parallel_tool_calls, + fast_mode = fast_mode, ): yield line return @@ -1296,6 +1508,7 @@ class ExternalProviderClient: stop: Optional[Union[str, list[str]]] = None, service_tier: Optional[str] = None, parallel_tool_calls: Optional[bool] = None, + fast_mode: Optional[bool] = None, ) -> AsyncGenerator[str, None]: """ Call the Anthropic Messages API and translate its SSE to OpenAI format. @@ -1415,6 +1628,11 @@ class ExternalProviderClient: "media_type": media_type, "data": b64data, }, + # Opt into Anthropic's natural-citation + # pipeline; without this no citations_delta + # events fire. See + # https://platform.claude.com/docs/en/build-with-claude/citations + "citations": {"enabled": True}, } if title: doc_block["title"] = title @@ -1426,6 +1644,7 @@ class ExternalProviderClient: "type": "url", "url": url, }, + "citations": {"enabled": True}, } if title: doc_block["title"] = title @@ -1655,25 +1874,19 @@ class ExternalProviderClient: ) body["tools"] = anthropic_tools - # Anthropic server-side web_fetch — see - # https://platform.claude.com/docs/en/agents-and-tools/tool-use/web-fetch-tool - # `web_fetch_20250910` reads a single URL (text or PDF) and - # returns a document block in a `web_fetch_tool_result`. For - # safety Anthropic only lets the model fetch URLs that already - # appeared in the conversation (user message, prior tool - # result, web_search hit) — there is no domain restriction we - # have to apply locally. No beta header is required today; the - # tool ships under the standard `2023-06-01` API version. We - # mirror the web_search wiring: max_uses cap, opt in via - # `enabled_tools=["web_fetch"]`, citations off by default - # because the frontend already paints source pills from the - # generic tool_end payload. + # Anthropic server-side web_fetch reads a single URL (text/PDF) + # and returns a `web_fetch_tool_result` document block. Opt in + # via `enabled_tools=["web_fetch"]`; no beta header required. + # `_anthropic_web_fetch_version` picks `web_fetch_20260209` + # (dynamic filtering) for Opus 4.6/4.7 + Sonnet 4.6, falling + # back to `web_fetch_20250910` elsewhere; mismatched variants + # return 400 so the per-model picker is required. web_fetch_enabled = bool(enabled_tools and "web_fetch" in enabled_tools) if web_fetch_enabled: anthropic_tools = list(body.get("tools") or []) anthropic_tools.append( { - "type": "web_fetch_20250910", + "type": _anthropic_web_fetch_version(model), "name": "web_fetch", "max_uses": 5, } @@ -1770,6 +1983,13 @@ class ExternalProviderClient: ] } + # fast_mode is Opus 4.6/4.7 only; silently drop elsewhere. + # Incompatible with the Priority service_tier (frontend gate + # prevents both at once; backend lets Anthropic 400 if combined). + fast_mode_active = bool(fast_mode) and _anthropic_supports_fast_mode(model) + if fast_mode_active: + body["speed"] = "fast" + url = f"{self.base_url}/messages" completion_id = f"chatcmpl-anthropic-{model.replace('/', '-')}" @@ -1828,6 +2048,8 @@ class ExternalProviderClient: beta_parts.append(_ANTHROPIC_CODE_EXECUTION_BETA) if compaction_active and _ANTHROPIC_COMPACTION_BETA not in beta_parts: beta_parts.append(_ANTHROPIC_COMPACTION_BETA) + if fast_mode_active and _ANTHROPIC_FAST_MODE_BETA not in beta_parts: + beta_parts.append(_ANTHROPIC_FAST_MODE_BETA) if beta_parts: request_headers["anthropic-beta"] = ",".join(beta_parts) @@ -1926,6 +2148,12 @@ class ExternalProviderClient: # the next turn. current_compaction: Optional[dict[str, Any]] = None compaction_blocks_seen = 0 + # Document citations from ``citations_delta`` events. + # Deduped by type-specific anchor key; inline [N] is + # injected after each cited run, and the full list is + # forwarded as a synthetic document_citations tool_event + # on message_stop for the Sources panel. + document_citations: list[dict[str, Any]] = [] # Counts surfaced in the final log line so reports of # "Code execution did nothing" can be triaged at a # glance. generated_files_count is interesting for the @@ -2277,10 +2505,27 @@ class ExternalProviderClient: thinking_open = False if text: yield _content_chunk(text) - # Citations on text deltas are attached - # per-call by Anthropic via the - # `web_search_tool_result` block; we don't - # need to scrape them off the text events. + # web_search citations: web_search_tool_result. + # User-doc citations: citations_delta below. + elif delta_type == "citations_delta": + # One citation per event; collapse onto a + # numbered footnote list and inject [N] + # inline. See + # https://platform.claude.com/docs/en/build-with-claude/citations + cit = delta.get("citation") + if isinstance(cit, dict): + key = _anthropic_citation_key(cit) + idx_for_marker: Optional[int] = None + for idx, existing in enumerate( + document_citations, start = 1 + ): + if existing.get("_key") == key: + idx_for_marker = idx + break + if idx_for_marker is None: + document_citations.append({**cit, "_key": key}) + idx_for_marker = len(document_citations) + yield _content_chunk(f"[{idx_for_marker}]") elif delta_type == "input_json_delta": # Streamed partial_json carrying tool inputs # — the search query for web_search, or the @@ -2569,6 +2814,29 @@ class ExternalProviderClient: # finish_reason="stop" chunk that would # truncate the rendered message in the UI. mapped = _finish_reason_map.get(stop_reason, "stop") + # Streaming refusal: emit a visible notice + # plus an out-of-band _toolEvent so the + # frontend can prune the refused turn. + # The mapped finish_reason is + # "content_filter" per OpenAI spec. + # https://platform.claude.com/docs/en/test-and-evaluate/strengthen-guardrails/handle-streaming-refusals + if stop_reason == "refusal": + logger.warning( + "Anthropic refusal stop_reason (model=%s)", + model, + ) + # Drop signal rides _toolEvent (not + # text) so assistant content cannot + # spoof a context reset. + yield _content_chunk( + "\n\n_The response was stopped by " + "Anthropic's safety classifier. Edit " + "or remove the previous turn and try " + "again._" + ) + yield _emit_tool_event( + {"type": "anthropic_refusal"} + ) if mapped is not None: chunk = { "id": completion_id, @@ -2587,6 +2855,29 @@ class ExternalProviderClient: if thinking_open: yield _content_chunk("") thinking_open = False + # Forward document_citations so the Sources + # panel can render the inline [N] footnotes. + # ``cited_text`` is truncated server-side to + # keep SSE bytes bounded on long spans. + if document_citations: + clean_cits = [] + for c in document_citations: + entry = {k: v for k, v in c.items() if k != "_key"} + cited = entry.get("cited_text") + if ( + isinstance(cited, str) + and len(cited) > _CITED_TEXT_MAX_LEN + ): + entry["cited_text"] = ( + cited[:_CITED_TEXT_MAX_LEN] + "…" + ) + clean_cits.append(entry) + yield _emit_tool_event( + { + "type": "document_citations", + "citations": clean_cits, + } + ) # Final include_usage-style chunk so callers can # see cache_creation / cache_read without # scraping the server log. @@ -3114,6 +3405,65 @@ class ExternalProviderClient: # see. latched_container_id: Optional[str] = None container_id_emitted = False + # Buffer for a citation marker straddling two delta events; + # prepended onto the next delta. See _split_pending_citation_tail. + pending_marker_tail: str = "" + # Segments deferred while their markers reference unseen + # source_ids; held in arrival order so output never + # leapfrogs an earlier deferred segment. Flushed on + # annotation events and force-flushed at end-of-stream + # with leftover private-use codepoints stripped. + pending_citation_segments: list[str] = [] + + def _drain_pending_segments(force: bool) -> str: + """Re-attempt resolution on buffered segments in order. + Stops at the first still-unresolved segment unless + ``force`` (end-of-stream), where lingering markers are stripped.""" + out: list[str] = [] + while pending_citation_segments: + seg = pending_citation_segments[0] + rewritten, unresolved = _rewrite_citation_markers_partial( + seg, + all_url_citations, + ) + if unresolved and not force: + pending_citation_segments[0] = rewritten + break + if unresolved and force: + rewritten = _replace_openai_citation_markers( + rewritten, + all_url_citations, + ) + pending_citation_segments.pop(0) + if rewritten: + out.append(rewritten) + return "".join(out) + + def _flush_pending_marker_tail(tail: str) -> str: + """Render any leftover citation tail at end-of-stream. + + Unterminated tails drop (no annotation to bind to). If the + close byte arrived concatenated, rewrite then scrub any + residual private-use bytes and any orphan ``cite`` + literal so the renderer never sees raw markup. url_citations + are aggregated separately and applied to web_search tool_end. + """ + if not tail: + return "" + if _OPENAI_CITE_STOP not in tail: + # Unterminated: drop the whole tail, otherwise the + # residual ``cite`` would leak as plain text. + return "" + rendered = _replace_openai_citation_markers( + tail, all_url_citations + ) + # Scrub residual private-use bytes (e.g. a partial opener). + for ch in ("", "", ""): + rendered = rendered.replace(ch, "") + # Drop any orphan ``cite`` literal -- meaningless + # without its closing byte and matching url_citation. + rendered = re.sub(r"^cite\S*", "", rendered) + return rendered def _emit_tool_event(payload: dict[str, Any]) -> str: chunk = { @@ -3171,16 +3521,35 @@ class ExternalProviderClient: def _record_url_citation(payload: dict[str, Any]) -> None: """Append a url_citation onto the shared all_url_citations - list. Dedup by URL — the same source can be cited multiple - times across deltas. We do NOT try to attribute citations - to individual web_search_call invocations because OpenAI's - annotation events don't carry that linkage.""" + list. Dedup by URL — the same URL can be cited many + times under different ``source_id`` aliases (one per + span/locator), so collect every alias we see onto + the matching entry's ``source_ids`` list. The + delta-text rewriter resolves any of those aliases + back to this entry's URL. The id may live under + ``source_id``, ``id``, or ``locator`` across the + Responses API revisions.""" if payload.get("type") != "url_citation": return url = payload.get("url", "") if not url: return - if any(c["url"] == url for c in all_url_citations): + source_id = ( + payload.get("source_id") + or payload.get("id") + or payload.get("locator") + or "" + ) + # Single pass: either backfill aliases onto an + # existing URL entry (and return) or fall through + # to append a fresh one. + for c in all_url_citations: + if c["url"] != url: + continue + if source_id: + aliases = c.setdefault("source_ids", []) + if source_id not in aliases: + aliases.append(source_id) return title = payload.get("title") or url snippet = payload.get("snippet") or payload.get("quote") or "" @@ -3189,6 +3558,7 @@ class ExternalProviderClient: "url": url, "title": title, "snippet": snippet, + "source_ids": [source_id] if source_id else [], } ) @@ -3245,6 +3615,28 @@ class ExternalProviderClient: if not data_str: continue if data_str == "[DONE]": + # Flush any held-over partial marker; strip + # private-use bytes so garbled glyphs don't leak. + if pending_marker_tail: + flushed = _flush_pending_marker_tail( + pending_marker_tail + ) + pending_marker_tail = "" + if flushed: + if reasoning_open: + yield _chunk_with_text("") + reasoning_open = False + yield _chunk_with_text(flushed) + # Force-drain any segment still awaiting an + # annotation; lingering codepoints are stripped. + tail_flushed = _drain_pending_segments( + force = True, + ) + if tail_flushed: + if reasoning_open: + yield _chunk_with_text("") + reasoning_open = False + yield _chunk_with_text(tail_flushed) if not done_emitted: yield "data: [DONE]" done_emitted = True @@ -3259,22 +3651,57 @@ class ExternalProviderClient: if event_type == "response.output_text.delta": delta_text = event.get("delta", "") - if delta_text: - if reasoning_open: - yield _chunk_with_text("") - reasoning_open = False - yield _chunk_with_text(delta_text) - # Some API versions inline url citations on the - # delta event itself rather than as a separate - # response.output_text.annotation.added event. + # Process inline annotations first so source_ids + # referenced by same-delta markers are in the lookup + # before the rewriter runs. Some API versions inline + # url citations on the delta event itself. for ann in event.get("annotations") or []: if isinstance(ann, dict): _record_url_citation(ann) + if delta_text or pending_marker_tail: + # Prepend any held-over tail so a marker + # straddling two SSE events resolves cleanly. + combined = pending_marker_tail + delta_text + head, pending_marker_tail = ( + _split_pending_citation_tail(combined) + ) + if head: + if reasoning_open: + yield _chunk_with_text("") + reasoning_open = False + # Re-attempt earlier deferred segments first + # so output stays in order; the needed + # annotation may have arrived inline above. + flushed = _drain_pending_segments( + force = False, + ) + if flushed: + yield _chunk_with_text(flushed) + head_rewritten, has_unresolved = ( + _rewrite_citation_markers_partial( + head, + all_url_citations, + ) + ) + if has_unresolved or pending_citation_segments: + pending_citation_segments.append( + head_rewritten + ) + elif head_rewritten: + yield _chunk_with_text(head_rewritten) elif event_type == "response.output_text.annotation.added": ann = event.get("annotation") if isinstance(ann, dict): _record_url_citation(ann) + flushed = _drain_pending_segments( + force = False, + ) + if flushed: + if reasoning_open: + yield _chunk_with_text("") + reasoning_open = False + yield _chunk_with_text(flushed) elif event_type == "response.output_item.added": # Track the call early but do NOT emit tool_start @@ -3513,6 +3940,34 @@ class ExternalProviderClient: ) if isinstance(completed_usage, dict): last_usage = completed_usage + # Flush any unterminated citation tail + # held over from the last delta. By + # the time we get here every annotation + # has been recorded so a late-arriving + # source_id may resolve cleanly; if it + # still doesn't, the helper strips the + # private-use bytes so no garbled + # glyph reaches the user. + if pending_marker_tail: + flushed = _flush_pending_marker_tail( + pending_marker_tail + ) + pending_marker_tail = "" + if flushed: + if reasoning_open: + yield _chunk_with_text("") + reasoning_open = False + yield _chunk_with_text(flushed) + # Force-drain any segment still awaiting an + # annotation; lingering codepoints are stripped. + tail_flushed = _drain_pending_segments( + force = True, + ) + if tail_flushed: + if reasoning_open: + yield _chunk_with_text("") + reasoning_open = False + yield _chunk_with_text(tail_flushed) if reasoning_open: yield _chunk_with_text("") reasoning_open = False @@ -3605,6 +4060,29 @@ class ExternalProviderClient: ) if isinstance(incomplete_usage, dict): last_usage = incomplete_usage + # Same flush as response.completed -- + # truncated streams can leave a half- + # marker in the buffer. + if pending_marker_tail: + flushed = _flush_pending_marker_tail( + pending_marker_tail + ) + pending_marker_tail = "" + if flushed: + if reasoning_open: + yield _chunk_with_text("") + reasoning_open = False + yield _chunk_with_text(flushed) + # Force-drain any segment still awaiting an + # annotation; lingering codepoints are stripped. + tail_flushed = _drain_pending_segments( + force = True, + ) + if tail_flushed: + if reasoning_open: + yield _chunk_with_text("") + reasoning_open = False + yield _chunk_with_text(tail_flushed) if reasoning_open: yield _chunk_with_text("") reasoning_open = False @@ -4110,6 +4588,17 @@ def _build_usage_chunk( "cache_creation_input_tokens": cache_creation, "cache_read_input_tokens": cache_read, } + # Forward 5m/1h cache-write breakdown so cost calc applies the + # 2x 1h premium instead of defaulting to 5m on chat-style. + cc_breakdown = last_usage.get("cache_creation") + if isinstance(cc_breakdown, dict) and cc_breakdown: + usage_block["cache_creation"] = cc_breakdown + # Propagate fast-mode `usage.speed` so the cost ledger can apply + # the 6x multiplier without re-derivation (Anthropic falls back + # to "standard" when fast-mode is unsupported or rate-limited). + speed = last_usage.get("speed") + if speed in ("fast", "standard"): + usage_block["speed"] = speed else: prompt_tokens = last_usage.get("input_tokens") or 0 cached = 0 diff --git a/studio/backend/core/inference/pricing.py b/studio/backend/core/inference/pricing.py index 74c57fa594..84eec17e84 100644 --- a/studio/backend/core/inference/pricing.py +++ b/studio/backend/core/inference/pricing.py @@ -1,50 +1,24 @@ # SPDX-License-Identifier: AGPL-3.0-only # Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0 -"""Static per-MTok pricing tables for external providers, plus a -``calculate_cost`` helper that turns an upstream ``usage`` block into -a USD figure for surfacing in the chat UI. +"""Static per-MTok pricing tables and ``calculate_cost`` helper for +turning an upstream ``usage`` block into a USD figure. -Neither the Anthropic Messages API nor the OpenAI Responses API -reports a ``cost`` field on the response. Both expose detailed token -counts (input, output, cache hits, server-tool invocations); pricing -multipliers live in the provider docs. We fold the docs into a static -table here, multiply by the usage block, and emit a per-turn cost + -running session total client-side. - -Sources (verified live 2026-05-22): -- Anthropic models overview: - https://platform.claude.com/docs/en/about-claude/models/overview -- Anthropic prompt-caching multipliers (5m write 1.25x, 1h write 2x, - read 0.1x): - https://platform.claude.com/docs/en/build-with-claude/prompt-caching -- Anthropic web search ($10 / 1000 searches, code execution - free-with-paid when paired with the newer web tools): - https://platform.claude.com/docs/en/agents-and-tools/tool-use/web-search-tool - https://platform.claude.com/docs/en/agents-and-tools/tool-use/code-execution-tool -- OpenAI pricing page (input / output per MTok per model family): - https://platform.openai.com/docs/pricing +Sources: Anthropic prompt-caching docs (5m write 1.25x, 1h write 2x, +read 0.1x), web search ($10/1000), code execution; OpenAI pricing page. """ from __future__ import annotations from typing import Any, Optional -# Per-million-token base pricing. `cache_5m_write_mult`, `cache_1h_write_mult`, -# `cache_read_mult` are multipliers ON `input_per_mtok` -- not absolute prices -- -# matching how Anthropic publishes them (5m write = 1.25x base, etc.). -# -# `input_per_mtok` and `output_per_mtok` are USD per 1,000,000 tokens. +# Per-MTok base pricing in USD. Cache multipliers are applied ON +# `input_per_mtok` (not absolute prices), matching Anthropic's docs. ANTHROPIC_PRICING: dict[str, dict[str, float]] = { "claude-opus-4-7": {"input_per_mtok": 5.0, "output_per_mtok": 25.0}, "claude-opus-4-6": {"input_per_mtok": 5.0, "output_per_mtok": 25.0}, - # Canonical 4.5 ids are referenced from backend defaults (e.g. - # PROVIDER_REGISTRY['anthropic'].default_models) without the date - # suffix. The dated ids ARE the canonical names per Anthropic's - # models overview, but lookups for the bare id ("claude-opus-4-5") - # don't prefix-match the dated key the other way around, so we - # alias both forms here. Otherwise calculate_cost returns - # priced=False + zero cost for the common ids. + # Alias both the bare id and dated id: backend defaults reference + # the bare form, which won't prefix-match the dated key. "claude-opus-4-5": {"input_per_mtok": 5.0, "output_per_mtok": 25.0}, "claude-opus-4-5-20251101": {"input_per_mtok": 5.0, "output_per_mtok": 25.0}, "claude-opus-4-1": {"input_per_mtok": 15.0, "output_per_mtok": 75.0}, @@ -59,19 +33,9 @@ ANTHROPIC_PRICING: dict[str, dict[str, float]] = { } OPENAI_PRICING: dict[str, dict[str, float]] = { - # All values verified against developers.openai.com/api/docs/pricing - # 2026-05-22. Update against the live pricing page on every model launch. - # Initial commit underbilled every gpt-5.x family 2-6x -- fixed here - # after PR review caught it via doc cross-check. - # - # `long_context_input_per_mtok` / `long_context_output_per_mtok` / - # `long_context_threshold` are populated when OpenAI publishes a - # second pricing tier for prompts above N input tokens. gpt-5.5 and - # gpt-5.4 cross over at 272k input tokens; the long-context rates - # are double the headline input price (and ~1.5x on output). Other - # families currently ship with a single rate (no `long_context_*` - # keys = no tier crossover). Reference: - # https://developers.openai.com/api/docs/pricing + # Verified against developers.openai.com/api/docs/pricing. + # `long_context_*` keys apply once input exceeds the threshold + # (gpt-5.5/5.4: 272k); families without these keys ship a single rate. "gpt-5.5": { "input_per_mtok": 5.0, "output_per_mtok": 30.0, @@ -91,43 +55,33 @@ OPENAI_PRICING: dict[str, dict[str, float]] = { "gpt-5.4-mini": {"input_per_mtok": 0.75, "output_per_mtok": 4.5}, "gpt-5.4-nano": {"input_per_mtok": 0.20, "output_per_mtok": 1.25}, "gpt-5.3-codex": {"input_per_mtok": 1.75, "output_per_mtok": 14.0}, - # chat-latest / gpt-5.3-chat-latest is an alias for the current - # ChatGPT model; same price as gpt-5.5. + # chat-latest aliases gpt-5.5. "gpt-5.3-chat-latest": {"input_per_mtok": 5.0, "output_per_mtok": 30.0}, "chat-latest": {"input_per_mtok": 5.0, "output_per_mtok": 30.0}, - # o-series and gpt-4.5: NOT currently listed on the pricing page. - # Removed to avoid silent-underbilling drift. Returning priced=False - # is honest; the UI can still render token counts. Restore with - # verified per-MTok rates if/when the page lists them again. + # o-series and gpt-4.5 are no longer on the pricing page; omit them + # so calculate_cost returns priced=False rather than silently $0. } # Shared multipliers (same across every Anthropic model). ANTHROPIC_CACHE_5M_WRITE_MULT = 1.25 ANTHROPIC_CACHE_1H_WRITE_MULT = 2.0 ANTHROPIC_CACHE_READ_MULT = 0.1 +# Anthropic fast-mode (Opus 4.6 / 4.7 only): 6x standard on input + output. +# https://platform.claude.com/docs/en/build-with-claude/fast-mode#pricing +ANTHROPIC_FAST_MODE_MULT = 6.0 -# OpenAI: cache reads are 0.1x base input, cache writes are not billed -# separately (the first prefix-write request just pays normal input). +# OpenAI: cache reads 0.1x; cache writes pay normal input price. OPENAI_CACHE_READ_MULT = 0.1 -# Server-tool surcharges. -# Anthropic: $10 / 1000 web searches; code_execution is $0.05/hr after -# 50 free hours/day per org (no per-org visibility here, so the -# calculator reports the marginal rate). +# Server-tool surcharges. Anthropic code_exec is $0.05/hr marginal +# (50 free hours/day per org, not visible here). ANTHROPIC_WEB_SEARCH_USD_PER_1K = 10.0 ANTHROPIC_CODE_EXEC_USD_PER_HOUR = 0.05 -# OpenAI: web_search is billed at $10/1000 calls plus the model's -# token rate for the returned search content (already captured under -# input/output_tokens). The hosted shell tool bills per 20-minute -# session per container memory tier (1g/4g/16g/64g at -# $0.03/$0.12/$0.48/$1.92). Since Studio doesn't surface the memory -# tier in the cost ledger and most users land on the default 1g, we -# bill the 1g rate ($0.09/hour) and let the user inspect the OpenAI -# dashboard for the exact figure on heavier configs. -# Source: developers.openai.com/api/docs/pricing 2026-05-22. +# OpenAI container bills per memory tier; we report the 1g default +# ($0.09/hour) since the tier isn't surfaced to the cost ledger. OPENAI_WEB_SEARCH_USD_PER_1K = 10.0 -OPENAI_CONTAINER_USD_PER_HOUR = 0.09 # 1g default tier; 3 x $0.03 / 60min +OPENAI_CONTAINER_USD_PER_HOUR = 0.09 # 1g default tier def _lookup(provider: str, model: str) -> Optional[dict[str, float]]: @@ -142,11 +96,13 @@ def _lookup(provider: str, model: str) -> Optional[dict[str, float]]: return None if model in table: return table[model] - # Fall back to a prefix match so date-suffixed snapshots - # ("gpt-5.5-2026-04-23") inherit the canonical-id prices. - for key, val in table.items(): - if model.startswith(key): - return val + # Longest-prefix match on a dash boundary: lets dated snapshots + # inherit canonical prices while preventing "claude-opus-4-15" + # from matching "claude-opus-4-1" or "gpt-5.5-prod" from matching + # "gpt-5.5-pro". Sort longest-first to pick the most specific row. + for key in sorted(table, key = len, reverse = True): + if model.startswith(key) and (len(model) == len(key) or model[len(key)] == "-"): + return table[key] return None @@ -155,28 +111,11 @@ def calculate_cost( model: str, usage: dict[str, Any], ) -> dict[str, float]: - """Return a per-turn USD cost breakdown. - - Returns a dict with the per-bucket cost AND the totals so the - frontend can render either a single number or a "where did the - money go" tooltip without re-doing the math: - - { - "input_usd": 0.0042, - "output_usd": 0.012, - "cache_write_usd": 0.0001, - "cache_read_usd": 0.0008, - "server_tools_usd": 0.01, - "total_usd": 0.0271, - "billable_input_tokens": 5023, # input + cache_create + cache_read - "billable_output_tokens": 480, - "model_priced": "claude-opus-4-7", - "priced": true, - } - - When the model isn't in the static table (new family, custom base - URL), `priced` is False and every USD field is 0.0; the frontend - can still show the token counts. + """Return a per-turn USD cost breakdown with per-bucket + total + fields so the frontend can render either a single number or a + tooltip without re-doing the math. When the model isn't in the + static table, ``priced`` is False and USD fields are 0.0 (token + counts still report). """ prices = _lookup(provider, model) out: dict[str, float] = { @@ -192,34 +131,64 @@ def calculate_cost( "priced": bool(prices), } - input_tokens = int(usage.get("input_tokens") or 0) - output_tokens = int(usage.get("output_tokens") or 0) - cache_creation = int(usage.get("cache_creation_input_tokens") or 0) - cache_read = int(usage.get("cache_read_input_tokens") or 0) - # OpenAI Responses reports cached tokens under input_tokens_details - # but ALSO folds them into the top-level input_tokens, so we don't - # add cache_read into the billable total again below (Anthropic - # excludes cache buckets from input_tokens, OpenAI includes them -- - # the two providers differ here and the calculator must match). - if provider == "openai": - details = usage.get("input_tokens_details") or {} + # Accept raw (input_tokens/output_tokens) and Studio chat-style + # (prompt_tokens/completion_tokens) envelopes. Cache buckets + # behave differently per envelope: + # raw Anthropic: input_tokens EXCLUDES cache buckets + # raw OpenAI: input_tokens INCLUDES cache_read + # Studio Anthropic: prompt_tokens INCLUDES cache_creation + cache_read + # Studio OpenAI: prompt_tokens == raw input_tokens + # Clamp tokens >=0 so corrupted payloads can't produce a negative bill. + cache_creation = max(0, int(usage.get("cache_creation_input_tokens") or 0)) + cache_read_native_present = ( + "cache_read_input_tokens" in usage + and usage.get("cache_read_input_tokens") is not None + ) + cache_read = max(0, int(usage.get("cache_read_input_tokens") or 0)) + # Fallback to mirrored prompt_tokens_details only when the native + # cache_read_input_tokens key is absent. An explicit native 0 is + # authoritative, so a stale mirrored block from a proxy can never + # inflate cache_read past the native count. + if not cache_read_native_present: + details = usage.get("prompt_tokens_details") or {} if isinstance(details, dict): - cache_read = max(cache_read, int(details.get("cached_tokens") or 0)) - # OpenAI: cache_read already counted inside input_tokens. + cache_read = max(0, int(details.get("cached_tokens") or 0)) + has_input_tokens = "input_tokens" in usage and usage.get("input_tokens") is not None + if has_input_tokens: + input_tokens = max(0, int(usage.get("input_tokens") or 0)) + else: + # Chat-style: peel cache buckets back out for Anthropic to + # recover the raw uncached prompt count. + prompt_tokens = max(0, int(usage.get("prompt_tokens") or 0)) + if provider == "anthropic": + input_tokens = max(0, prompt_tokens - cache_creation - cache_read) + else: + input_tokens = prompt_tokens + # Prefer raw output_tokens even when 0 (an `or` fallback would + # silently pick a stale completion_tokens). + if "output_tokens" in usage and usage.get("output_tokens") is not None: + output_tokens = max(0, int(usage.get("output_tokens") or 0)) + else: + output_tokens = max(0, int(usage.get("completion_tokens") or 0)) + if provider == "openai": + # Cached tokens land on either input_tokens_details (raw + # Responses) or prompt_tokens_details (Studio chat-style). + for key in ("input_tokens_details", "prompt_tokens_details"): + details = usage.get(key) or {} + if isinstance(details, dict): + cache_read = max(cache_read, int(details.get("cached_tokens") or 0)) + # OpenAI input_tokens already counts cache_read. out["billable_input_tokens"] = input_tokens + cache_creation else: - # Anthropic: input_tokens excludes cache_* buckets, add them all. + # Anthropic input_tokens excludes cache buckets; add them back. out["billable_input_tokens"] = input_tokens + cache_creation + cache_read out["billable_output_tokens"] = output_tokens if not prices: return out - # Long-context tier crossover (gpt-5.5 / gpt-5.4 today). OpenAI - # bills the whole turn at the long-context rate once the prompt - # crosses the threshold, NOT a per-token blend, so we pick a - # single (base, out_per) pair for this turn based on - # billable_input_tokens. + # Long-context tier: whole-turn flip (not per-token blend) once + # billable_input_tokens crosses the threshold. lc_thresh = prices.get("long_context_threshold") in_long_context_tier = ( lc_thresh is not None @@ -235,17 +204,27 @@ def calculate_cost( base = prices["input_per_mtok"] out_per = prices["output_per_mtok"] + # Anthropic fast-mode: 6x on input + output. Cache multipliers stack + # on top of fast-mode, so applying once to (base, out_per) propagates + # into the cache_*_usd buckets computed below. + if provider == "anthropic" and usage.get("speed") == "fast": + base *= ANTHROPIC_FAST_MODE_MULT + out_per *= ANTHROPIC_FAST_MODE_MULT + if out["model_priced"]: + out["model_priced"] = f"{out['model_priced']} (fast)" + out["input_usd"] = (input_tokens / 1_000_000.0) * base out["output_usd"] = (output_tokens / 1_000_000.0) * out_per if provider == "anthropic": - # Split cache_creation across 5m / 1h buckets when the - # response surfaces the breakdown. - cc_breakdown = usage.get("cache_creation") or {} - cc_5m = int(cc_breakdown.get("ephemeral_5m_input_tokens") or 0) - cc_1h = int(cc_breakdown.get("ephemeral_1h_input_tokens") or 0) + # Split cache_creation into 5m / 1h buckets when surfaced. + # Tolerate non-dict (some proxies fold to an int total). + cc_raw = usage.get("cache_creation") + cc_breakdown = cc_raw if isinstance(cc_raw, dict) else {} + cc_5m = max(0, int(cc_breakdown.get("ephemeral_5m_input_tokens") or 0)) + cc_1h = max(0, int(cc_breakdown.get("ephemeral_1h_input_tokens") or 0)) if cc_5m + cc_1h == 0 and cache_creation > 0: - # Fall back: assume default 5m pool when no breakdown is given. + # No breakdown -- assume default 5m pool. cc_5m = cache_creation out["cache_write_usd"] = ( cc_5m / 1_000_000.0 @@ -265,24 +244,17 @@ def calculate_cost( + code_exec_hours * ANTHROPIC_CODE_EXEC_USD_PER_HOUR ) else: - # OpenAI: cache writes share the base input price (no premium). - # Only cache reads get the 0.1x multiplier; subtract those from - # the input_usd we already counted so we don't double-bill. - # Anthropic excludes cache buckets from input_tokens, but - # OpenAI folds them in, so the math differs. + # OpenAI: cache writes pay base input; only cache reads get + # 0.1x. Subtract cached from already-counted input_usd to + # avoid double-billing (OpenAI folds cache into input_tokens). if cache_read > 0: non_cached_input = max(0, input_tokens - cache_read) out["input_usd"] = (non_cached_input / 1_000_000.0) * base out["cache_read_usd"] = ( (cache_read / 1_000_000.0) * base * OPENAI_CACHE_READ_MULT ) - # Server-tool surcharges. OpenAI doesn't include these on its - # `usage` object directly -- web_search invocations are counted - # from `ResponseFunctionWebSearch` items in the output array, - # and container hours come from the SSE translator's shell-tool - # accounting. Studio surfaces both under a normalised - # `openai_tool_use` key on the usage dict the SSE finaliser - # hands to this calculator. + # OpenAI server-tool surcharges arrive under `openai_tool_use` + # (normalised by the SSE finaliser from output array items). srv = usage.get("openai_tool_use") or {} if isinstance(srv, dict): web_searches = int(srv.get("web_search_requests") or 0) @@ -304,17 +276,14 @@ def calculate_cost( def pricing_snapshot() -> dict[str, Any]: - """Whole pricing table, for the /api/providers/pricing endpoint. - - Returns a flat structure the frontend can hand to its cost - formatter without re-implementing the multipliers. - """ + """Whole pricing table for the /api/providers/pricing endpoint.""" return { "anthropic": { "models": dict(ANTHROPIC_PRICING), "cache_5m_write_mult": ANTHROPIC_CACHE_5M_WRITE_MULT, "cache_1h_write_mult": ANTHROPIC_CACHE_1H_WRITE_MULT, "cache_read_mult": ANTHROPIC_CACHE_READ_MULT, + "fast_mode_mult": ANTHROPIC_FAST_MODE_MULT, "web_search_usd_per_1k": ANTHROPIC_WEB_SEARCH_USD_PER_1K, "code_execution_usd_per_hour": ANTHROPIC_CODE_EXEC_USD_PER_HOUR, }, diff --git a/studio/backend/main.py b/studio/backend/main.py index 004ae404cd..689241b915 100644 --- a/studio/backend/main.py +++ b/studio/backend/main.py @@ -49,6 +49,8 @@ import shutil import warnings from contextlib import asynccontextmanager from importlib.metadata import PackageNotFoundError, version as package_version +from typing import Optional +from urllib.parse import urlparse _STUDIO_INSTALL_ID_RE = _re.compile(r"^[0-9a-f]{64}$") @@ -715,10 +717,8 @@ def _strip_crossorigin(html_bytes: bytes) -> bytes: def _inject_bootstrap(html_bytes: bytes, app: FastAPI): """Inject bootstrap credentials when password change is pending. - - Returns ``(html_bytes, script_nonce_or_None)``. Callers must forward - the nonce via ``_CSP_SCRIPT_NONCE_HEADER`` so the inline script is - not blocked by CSP. + Returns ``(html_bytes, script_nonce_or_None)``; callers forward the + nonce via ``_CSP_SCRIPT_NONCE_HEADER`` so CSP allows the inline script. """ import json as _json import secrets as _secrets @@ -743,6 +743,86 @@ def _inject_bootstrap(html_bytes: bytes, app: FastAPI): return html.encode("utf-8"), nonce +_DEFAULT_PORTS = {"http": 80, "https": 443, "ws": 80, "wss": 443} + + +def _canonical_origin(scheme: str, netloc: str) -> Optional[tuple[str, str, int]]: + """Canonicalise an Origin to ``(scheme, host, port)`` for equality. + Browsers strip default ports (RFC 6454 sec 6.1) and scheme/host are + case-insensitive (RFC 3986), so bare string compare misclassifies + same-origin requests as cross-origin. Returns ``None`` on unparseable + input so callers fall to the safer cross-origin default. + """ + scheme = (scheme or "").strip().lower() + if not scheme or not netloc: + return None + # Strip userinfo (RFC 3986); Origin never carries credentials. + if "@" in netloc: + netloc = netloc.rsplit("@", 1)[1] + # IPv6 hosts use brackets (RFC 3986 sec 3.2.2): ``[::1]:8902``. Bare + # ``partition(":")`` mis-parses these and breaks ``unsloth studio -H ::1``. + if netloc.startswith("["): + close = netloc.find("]") + if close == -1: + return None + host = netloc[1:close] + rest = netloc[close + 1 :] + if rest.startswith(":"): + port_str = rest[1:] + elif rest == "": + port_str = "" + else: + return None + else: + host, _, port_str = netloc.partition(":") + host = host.strip().lower() + if not host: + return None + if port_str: + try: + port = int(port_str) + except ValueError: + return None + else: + port = _DEFAULT_PORTS.get(scheme, 0) + return (scheme, host, port) + + +def _is_same_origin_request(request: Request) -> bool: + """True when Origin is missing or matches request's scheme://host:port. + Top-level same-document GETs omit Origin, so missing counts as same-origin. + Callers must also emit ``Vary: Origin``. Both sides are canonicalised via + :func:`_canonical_origin` so default-port stripping and scheme/host case + do not misclassify same-origin requests as cross-origin. + """ + origin = request.headers.get("origin") + if origin is None: + # Missing header: top-level same-document GETs omit Origin. + return True + # Empty string is not a valid serialised origin (RFC 6454 sec 6.1). + if not origin: + return False + # "null" token (sandboxed iframes, file:// pages) is never same-origin. + if origin == "null": + return False + # ``urlparse`` raises ``ValueError`` on malformed IPv6 brackets; swallow + # so a garbage Origin doesn't 500 the SPA handler. + try: + parsed = urlparse(origin) + except ValueError: + return False + origin_canon = _canonical_origin(parsed.scheme, parsed.netloc) + if origin_canon is None: + return False + try: + self_canon = _canonical_origin(request.url.scheme, request.url.netloc) + except ValueError: + return False + if self_canon is None: + return False + return origin_canon == self_canon + + def setup_frontend(app: FastAPI, build_path: Path): """Mount frontend static files (optional)""" if not build_path.exists(): @@ -753,11 +833,18 @@ def setup_frontend(app: FastAPI, build_path: Path): if assets_dir.exists(): app.mount("/assets", StaticFiles(directory = assets_dir), name = "assets") - def _build_index_response() -> Response: + def _build_index_response(request: Request) -> Response: content = (build_path / "index.html").read_bytes() content = _strip_crossorigin(content) - content, nonce = _inject_bootstrap(content, app) - headers = {"Cache-Control": "no-cache, no-store, must-revalidate"} + # Bootstrap pw is same-origin only; Vary: Origin keeps caches honest. + if _is_same_origin_request(request): + content, nonce = _inject_bootstrap(content, app) + else: + nonce = None + headers = { + "Cache-Control": "no-cache, no-store, must-revalidate", + "Vary": "Origin", + } if nonce: headers[_CSP_SCRIPT_NONCE_HEADER] = nonce return Response( @@ -767,11 +854,11 @@ def setup_frontend(app: FastAPI, build_path: Path): ) @app.get("/") - async def serve_root(): - return _build_index_response() + async def serve_root(request: Request): + return _build_index_response(request) @app.get("/{full_path:path}") - async def serve_frontend(full_path: str): + async def serve_frontend(request: Request, full_path: str): if full_path in {"api", "v1"} or full_path.startswith(("api/", "v1/")): return {"error": "API endpoint not found"} @@ -785,6 +872,6 @@ def setup_frontend(app: FastAPI, build_path: Path): return FileResponse(file_path) # Serve index.html as bytes — avoids Content-Length mismatch - return _build_index_response() + return _build_index_response(request) return True diff --git a/studio/backend/models/inference.py b/studio/backend/models/inference.py index 0c1c3f814b..040bb6ce6c 100644 --- a/studio/backend/models/inference.py +++ b/studio/backend/models/inference.py @@ -828,6 +828,16 @@ class ChatCompletionRequest(BaseModel): "default (which is `true` everywhere today)." ), ) + fast_mode: Optional[bool] = Field( + None, + description = ( + "[x-unsloth] Anthropic fast-mode toggle. On Claude Opus 4.6 / " + "4.7 adds the `fast-mode-2026-02-01` beta header and sends " + "`speed: 'fast'` for higher OTPS at premium pricing. Silently " + "ignored on every other model + provider. See " + "https://platform.claude.com/docs/en/build-with-claude/fast-mode" + ), + ) @model_validator(mode = "after") def _resolve_missing_tool_call_ids(self) -> "ChatCompletionRequest": diff --git a/studio/backend/routes/inference.py b/studio/backend/routes/inference.py index c14d721cec..72ab4afed8 100644 --- a/studio/backend/routes/inference.py +++ b/studio/backend/routes/inference.py @@ -1897,6 +1897,7 @@ async def _proxy_to_external_provider( stop = payload.stop, service_tier = payload.service_tier, parallel_tool_calls = payload.parallel_tool_calls, + fast_mode = payload.fast_mode, stream = payload.stream, ) try: diff --git a/studio/backend/tests/test_anthropic_citations.py b/studio/backend/tests/test_anthropic_citations.py new file mode 100644 index 0000000000..ab5ba10b56 --- /dev/null +++ b/studio/backend/tests/test_anthropic_citations.py @@ -0,0 +1,353 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0 + +"""Tests for Anthropic ``citations_delta`` handling in the streaming proxy. + +Verifies the proxy injects inline ``[N]`` markers after cited text, +dedupes by type-specific anchor (char_location, page_location, +content_block_location, search_result_location), forwards a synthetic +``document_citations`` tool_event at message_stop, and stays inert when +no citations_delta events fire. See +https://platform.claude.com/docs/en/build-with-claude/citations +""" + +import asyncio +import json + +import httpx + +from core.inference import external_provider as ep_mod +from core.inference.external_provider import ExternalProviderClient + + +def _drive(coro): + return asyncio.new_event_loop().run_until_complete(coro) + + +def _make_client() -> ExternalProviderClient: + return ExternalProviderClient( + provider_type = "anthropic", + base_url = "https://api.anthropic.com/v1", + api_key = "sk-ant-test", + ) + + +def _sse(events: list[dict]) -> bytes: + out = [] + for e in events: + ev = e.get("type", "message") + out.append(f"event: {ev}\ndata: {json.dumps(e)}\n\n") + return "".join(out).encode("utf-8") + + +def _capture(monkeypatch, events: list[dict]) -> list[str]: + def handler(request: httpx.Request) -> httpx.Response: + return httpx.Response( + 200, + content = _sse(events), + headers = {"content-type": "text/event-stream"}, + ) + + monkeypatch.setattr( + ep_mod, + "_http_client", + httpx.AsyncClient(transport = httpx.MockTransport(handler)), + ) + + lines: list[str] = [] + + async def run(): + client = _make_client() + try: + async for line in client.stream_chat_completion( + messages = [{"role": "user", "content": "what color is grass?"}], + model = "claude-opus-4-7", + max_tokens = 64, + ): + lines.append(line) + finally: + await client.close() + + _drive(run()) + return lines + + +def _message_start() -> dict: + return { + "type": "message_start", + "message": { + "id": "m1", + "content": [], + "model": "claude-opus-4-7", + "role": "assistant", + "stop_reason": None, + "usage": {"input_tokens": 5, "output_tokens": 2}, + }, + } + + +def _content_block_start_text() -> dict: + return { + "type": "content_block_start", + "index": 0, + "content_block": {"type": "text", "text": ""}, + } + + +def _text_delta(text: str, index: int = 0) -> dict: + return { + "type": "content_block_delta", + "index": index, + "delta": {"type": "text_delta", "text": text}, + } + + +def _citations_delta(citation: dict, index: int = 0) -> dict: + return { + "type": "content_block_delta", + "index": index, + "delta": {"type": "citations_delta", "citation": citation}, + } + + +def _content_block_stop(index: int = 0) -> dict: + return {"type": "content_block_stop", "index": index} + + +def _message_delta_end() -> dict: + return {"type": "message_delta", "delta": {"stop_reason": "end_turn"}} + + +def _message_stop() -> dict: + return {"type": "message_stop"} + + +def _joined(lines: list[str]) -> str: + return "\n".join(lines) + + +def test_no_citations_stream_unchanged(monkeypatch): + """Plain text streams pass through with no inline markers and no + document_citations tool_event.""" + lines = _capture( + monkeypatch, + [ + _message_start(), + _content_block_start_text(), + _text_delta("Grass is green."), + _content_block_stop(), + _message_delta_end(), + _message_stop(), + ], + ) + body = _joined(lines) + assert "Grass is green." in body + assert "document_citations" not in body + assert "[1]" not in body + + +def test_single_char_location_emits_inline_marker(monkeypatch): + cit = { + "type": "char_location", + "cited_text": "The grass is green.", + "document_index": 0, + "document_title": "Example", + "start_char_index": 0, + "end_char_index": 20, + } + lines = _capture( + monkeypatch, + [ + _message_start(), + _content_block_start_text(), + _text_delta("Grass is green."), + _citations_delta(cit), + _content_block_stop(), + _message_delta_end(), + _message_stop(), + ], + ) + body = _joined(lines) + assert "Grass is green." in body + assert "[1]" in body, body + assert "document_citations" in body, body + assert '"document_index": 0' in body, body + assert "_key" not in body, body + + +def test_duplicate_citation_dedupes_to_same_number(monkeypatch): + cit = { + "type": "char_location", + "document_index": 0, + "document_title": "Example", + "start_char_index": 0, + "end_char_index": 20, + "cited_text": "The grass is green.", + } + lines = _capture( + monkeypatch, + [ + _message_start(), + _content_block_start_text(), + _text_delta("Grass."), + _citations_delta(cit), + _text_delta(" Still green."), + _citations_delta(cit), + _content_block_stop(), + _message_delta_end(), + _message_stop(), + ], + ) + body = _joined(lines) + assert body.count("[1]") == 2, body + citation_blob = body[body.index("document_citations") :] + assert citation_blob.count('"start_char_index"') == 1, citation_blob + + +def test_distinct_sources_get_distinct_numbers(monkeypatch): + cit1 = { + "type": "char_location", + "document_index": 0, + "document_title": "Doc A", + "start_char_index": 0, + "end_char_index": 5, + } + cit2 = { + "type": "page_location", + "document_index": 1, + "document_title": "Doc B", + "start_page_number": 3, + "end_page_number": 4, + } + cit3 = { + "type": "content_block_location", + "document_index": 2, + "document_title": "Doc C", + "start_block_index": 0, + "end_block_index": 1, + } + lines = _capture( + monkeypatch, + [ + _message_start(), + _content_block_start_text(), + _text_delta("First"), + _citations_delta(cit1), + _text_delta(" Second"), + _citations_delta(cit2), + _text_delta(" Third"), + _citations_delta(cit3), + _content_block_stop(), + _message_delta_end(), + _message_stop(), + ], + ) + body = _joined(lines) + assert "[1]" in body and "[2]" in body and "[3]" in body, body + assert body.index("[1]") < body.index("[2]") < body.index("[3]") + + +def test_search_result_location_supported(monkeypatch): + cit = { + "type": "search_result_location", + "document_index": 0, + "document_title": "Anthropic Search Results", + "source": "https://example.com/doc.html", + "start_block_index": 0, + "end_block_index": 1, + "cited_text": "blah", + } + lines = _capture( + monkeypatch, + [ + _message_start(), + _content_block_start_text(), + _text_delta("Some sourced fact."), + _citations_delta(cit), + _content_block_stop(), + _message_delta_end(), + _message_stop(), + ], + ) + body = _joined(lines) + assert "[1]" in body + assert "search_result_location" in body + assert "example.com/doc.html" in body + + +def test_same_start_different_end_offsets_get_distinct_numbers(monkeypatch): + """Same start_char_index + different end_char_index = distinct spans, + so they must get distinct footnote numbers (ranges use exclusive end).""" + cit_a = { + "type": "char_location", + "document_index": 0, + "document_title": "Doc", + "start_char_index": 100, + "end_char_index": 150, + "cited_text": "first half", + } + cit_b = { + "type": "char_location", + "document_index": 0, + "document_title": "Doc", + "start_char_index": 100, + "end_char_index": 250, + "cited_text": "wider span", + } + lines = _capture( + monkeypatch, + [ + _message_start(), + _content_block_start_text(), + _text_delta("A "), + _citations_delta(cit_a), + _text_delta(" and B "), + _citations_delta(cit_b), + _content_block_stop(), + _message_delta_end(), + _message_stop(), + ], + ) + body = _joined(lines) + assert "[1]" in body, body + assert "[2]" in body, body + + +def test_search_result_location_different_indices_get_distinct_numbers(monkeypatch): + """Same source + different search_result_index = distinct footnotes + (matches the Anthropic search-result citation contract).""" + cit_a = { + "type": "search_result_location", + "search_result_index": 0, + "source": "https://example.com/result.html", + "title": "Result", + "start_block_index": 0, + "end_block_index": 1, + "cited_text": "first", + } + cit_b = { + "type": "search_result_location", + "search_result_index": 1, + "source": "https://example.com/result.html", + "title": "Result", + "start_block_index": 0, + "end_block_index": 1, + "cited_text": "second", + } + lines = _capture( + monkeypatch, + [ + _message_start(), + _content_block_start_text(), + _text_delta("A "), + _citations_delta(cit_a), + _text_delta(" and B "), + _citations_delta(cit_b), + _content_block_stop(), + _message_delta_end(), + _message_stop(), + ], + ) + body = _joined(lines) + assert "[1]" in body, body + assert "[2]" in body, body diff --git a/studio/backend/tests/test_anthropic_citations_edge.py b/studio/backend/tests/test_anthropic_citations_edge.py new file mode 100644 index 0000000000..be1b5f7922 --- /dev/null +++ b/studio/backend/tests/test_anthropic_citations_edge.py @@ -0,0 +1,690 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0 + +"""Edge-case tests for Anthropic ``citations_delta`` handling. + +Complements ``test_anthropic_citations.py``. Covers malformed payloads, +unusual orderings, mixed citation types, and the ``citations: +{enabled: true}`` opt-in attached to translated ``input_document`` +blocks. See +https://platform.claude.com/docs/en/build-with-claude/citations and +https://platform.claude.com/docs/en/build-with-claude/search-results. +""" + +import asyncio +import json + +import httpx + +from core.inference import external_provider as ep_mod +from core.inference.external_provider import ExternalProviderClient + + +# ── shared SSE harness ─────────────────────────────────────── + + +def _drive(coro): + return asyncio.new_event_loop().run_until_complete(coro) + + +def _make_client() -> ExternalProviderClient: + return ExternalProviderClient( + provider_type = "anthropic", + base_url = "https://api.anthropic.com/v1", + api_key = "sk-ant-test", + ) + + +def _sse(events: list[dict]) -> bytes: + out = [] + for e in events: + ev = e.get("type", "message") + out.append(f"event: {ev}\ndata: {json.dumps(e)}\n\n") + return "".join(out).encode("utf-8") + + +def _capture( + monkeypatch, + events: list[dict], + *, + messages: list[dict] | None = None, + captured_body: dict | None = None, +) -> list[str]: + """Drive ``stream_chat_completion`` against a mocked Anthropic + response and return the SSE lines. Pass ``captured_body`` to also + capture the outgoing request body for assertions on the translated + Anthropic shape. + """ + + def handler(request: httpx.Request) -> httpx.Response: + if captured_body is not None: + try: + captured_body.update(json.loads(request.content.decode("utf-8"))) + except Exception: # pragma: no cover -- diagnostic only + pass + return httpx.Response( + 200, + content = _sse(events), + headers = {"content-type": "text/event-stream"}, + ) + + monkeypatch.setattr( + ep_mod, + "_http_client", + httpx.AsyncClient(transport = httpx.MockTransport(handler)), + ) + + lines: list[str] = [] + + async def run(): + client = _make_client() + try: + async for line in client.stream_chat_completion( + messages = messages + or [{"role": "user", "content": "what color is grass?"}], + model = "claude-opus-4-7", + max_tokens = 64, + ): + lines.append(line) + finally: + await client.close() + + _drive(run()) + return lines + + +def _message_start() -> dict: + return { + "type": "message_start", + "message": { + "id": "m1", + "content": [], + "model": "claude-opus-4-7", + "role": "assistant", + "stop_reason": None, + "usage": {"input_tokens": 5, "output_tokens": 2}, + }, + } + + +def _content_block_start_text() -> dict: + return { + "type": "content_block_start", + "index": 0, + "content_block": {"type": "text", "text": ""}, + } + + +def _text_delta(text: str, index: int = 0) -> dict: + return { + "type": "content_block_delta", + "index": index, + "delta": {"type": "text_delta", "text": text}, + } + + +def _citations_delta(citation: dict, index: int = 0) -> dict: + return { + "type": "content_block_delta", + "index": index, + "delta": {"type": "citations_delta", "citation": citation}, + } + + +def _content_block_stop(index: int = 0) -> dict: + return {"type": "content_block_stop", "index": index} + + +def _message_delta_end() -> dict: + return {"type": "message_delta", "delta": {"stop_reason": "end_turn"}} + + +def _message_stop() -> dict: + return {"type": "message_stop"} + + +def _joined(lines: list[str]) -> str: + return "\n".join(lines) + + +def _citation_payload(body: str) -> dict: + """Pull the ``document_citations`` synthetic tool_event from the + SSE body and return its payload. Raises if absent.""" + assert "document_citations" in body, body + for line in body.splitlines(): + if not line.startswith("data: "): + continue + try: + payload = json.loads(line[len("data: ") :]) + except json.JSONDecodeError: + continue + tool_event = payload.get("_toolEvent") if isinstance(payload, dict) else None + if ( + isinstance(tool_event, dict) + and tool_event.get("type") == "document_citations" + ): + return tool_event + raise AssertionError("document_citations event not parsed out of SSE body") + + +# ── edge cases ─────────────────────────────────────────────── + + +def test_citation_with_no_preceding_text_still_emits_marker(monkeypatch): + """citations_delta before any text_delta must not crash; marker + lands at the start of the block.""" + cit = { + "type": "char_location", + "document_index": 0, + "document_title": "X", + "start_char_index": 0, + "end_char_index": 5, + "cited_text": "x", + } + lines = _capture( + monkeypatch, + [ + _message_start(), + _content_block_start_text(), + _citations_delta(cit), + _text_delta("hello"), + _content_block_stop(), + _message_delta_end(), + _message_stop(), + ], + ) + body = _joined(lines) + assert "[1]" in body, body + assert "document_citations" in body, body + + +def test_citations_delta_with_non_dict_citation_is_ignored(monkeypatch): + """Non-dict ``delta.citation`` must not crash, emit a marker, or + poison the document_citations list.""" + lines = _capture( + monkeypatch, + [ + _message_start(), + _content_block_start_text(), + _text_delta("Hello."), + { + "type": "content_block_delta", + "index": 0, + "delta": {"type": "citations_delta", "citation": "not-a-dict"}, + }, + _content_block_stop(), + _message_delta_end(), + _message_stop(), + ], + ) + body = _joined(lines) + assert "Hello." in body + assert "[1]" not in body + assert "document_citations" not in body + + +def test_citations_delta_with_missing_citation_field_is_ignored(monkeypatch): + """Missing ``citation`` field is treated like a non-dict citation: + skip without crashing.""" + lines = _capture( + monkeypatch, + [ + _message_start(), + _content_block_start_text(), + _text_delta("Hello."), + { + "type": "content_block_delta", + "index": 0, + "delta": {"type": "citations_delta"}, + }, + _content_block_stop(), + _message_delta_end(), + _message_stop(), + ], + ) + body = _joined(lines) + assert "Hello." in body + assert "[1]" not in body + assert "document_citations" not in body + + +def test_char_location_with_reversed_indices_does_not_crash(monkeypatch): + """Malformed char_location with reversed indices must not crash; + the dedup key accepts any int pair and still surfaces a footnote.""" + cit = { + "type": "char_location", + "document_index": 0, + "document_title": "Doc", + "start_char_index": 300, + "end_char_index": 50, + "cited_text": "?", + } + lines = _capture( + monkeypatch, + [ + _message_start(), + _content_block_start_text(), + _text_delta("Weird."), + _citations_delta(cit), + _content_block_stop(), + _message_delta_end(), + _message_stop(), + ], + ) + body = _joined(lines) + assert "[1]" in body, body + payload = _citation_payload(body) + assert payload["citations"][0]["start_char_index"] == 300 + assert payload["citations"][0]["end_char_index"] == 50 + + +def test_page_location_missing_document_index_does_not_crash(monkeypatch): + """page_location missing ``document_index`` still produces a + footnote; dedup key falls back to ``None`` for the missing field.""" + cit = { + "type": "page_location", + "document_title": "Untitled PDF", + "start_page_number": 1, + "end_page_number": 2, + "cited_text": "p1", + } + lines = _capture( + monkeypatch, + [ + _message_start(), + _content_block_start_text(), + _text_delta("From the PDF:"), + _citations_delta(cit), + _content_block_stop(), + _message_delta_end(), + _message_stop(), + ], + ) + body = _joined(lines) + assert "[1]" in body, body + payload = _citation_payload(body) + assert payload["citations"][0].get("document_index") is None + + +def test_content_block_location_with_non_int_block_index_does_not_crash(monkeypatch): + """content_block_location with string block indices must not crash; + dedup key tolerates non-int values.""" + cit = { + "type": "content_block_location", + "document_index": 0, + "document_title": "Custom", + "start_block_index": "0", + "end_block_index": "1", + "cited_text": "anything", + } + lines = _capture( + monkeypatch, + [ + _message_start(), + _content_block_start_text(), + _text_delta("Cite."), + _citations_delta(cit), + _content_block_stop(), + _message_delta_end(), + _message_stop(), + ], + ) + body = _joined(lines) + assert "[1]" in body, body + payload = _citation_payload(body) + assert payload["citations"][0]["start_block_index"] == "0" + + +def test_unknown_citation_type_falls_back_to_stringified_key(monkeypatch): + """Unknown citation ``type`` (forward-compat) still dedupes: + identical ones collapse, differing ones get distinct numbers.""" + cit_a = { + "type": "future_shape_location", + "anchor": "abc", + "cited_text": "blah", + } + cit_b = { + "type": "future_shape_location", + "anchor": "xyz", + "cited_text": "blah", + } + lines = _capture( + monkeypatch, + [ + _message_start(), + _content_block_start_text(), + _text_delta("A"), + _citations_delta(cit_a), + _text_delta(" again"), + _citations_delta(cit_a), + _text_delta(" B"), + _citations_delta(cit_b), + _content_block_stop(), + _message_delta_end(), + _message_stop(), + ], + ) + body = _joined(lines) + # cit_a dedupes onto [1], cit_b gets [2]. + assert body.count("[1]") == 2, body + assert body.count("[2]") == 1, body + payload = _citation_payload(body) + assert len(payload["citations"]) == 2 + + +def test_mixed_citation_types_same_document_get_distinct_keys(monkeypatch): + """char_location and page_location on the same document_index are + distinct shapes; dedup key uses citation type as its first slot.""" + cit_char = { + "type": "char_location", + "document_index": 0, + "document_title": "Doc", + "start_char_index": 0, + "end_char_index": 10, + } + cit_page = { + "type": "page_location", + "document_index": 0, + "document_title": "Doc", + "start_page_number": 1, + "end_page_number": 2, + } + lines = _capture( + monkeypatch, + [ + _message_start(), + _content_block_start_text(), + _text_delta("char-cite"), + _citations_delta(cit_char), + _text_delta(" page-cite"), + _citations_delta(cit_page), + _content_block_stop(), + _message_delta_end(), + _message_stop(), + ], + ) + body = _joined(lines) + assert "[1]" in body and "[2]" in body, body + payload = _citation_payload(body) + assert len(payload["citations"]) == 2 + + +def test_cited_text_is_preserved_in_synthetic_event(monkeypatch): + """``cited_text`` must survive into the synthetic event so the + Sources panel can render it as a tooltip. Anthropic does not bill + cited_text against output tokens, so preserving it is free.""" + cit = { + "type": "char_location", + "document_index": 0, + "document_title": "Trustworthy Doc", + "start_char_index": 0, + "end_char_index": 20, + "cited_text": "The grass is green.", + } + lines = _capture( + monkeypatch, + [ + _message_start(), + _content_block_start_text(), + _text_delta("Grass is green."), + _citations_delta(cit), + _content_block_stop(), + _message_delta_end(), + _message_stop(), + ], + ) + body = _joined(lines) + payload = _citation_payload(body) + assert payload["citations"][0]["cited_text"] == "The grass is green." + + +def test_internal_key_field_never_leaks_to_client(monkeypatch): + """The internal ``_key`` dedup sentinel must be stripped before + the synthetic event is forwarded; it is not an Anthropic field.""" + cit = { + "type": "char_location", + "document_index": 0, + "document_title": "Doc", + "start_char_index": 0, + "end_char_index": 5, + "cited_text": "..", + } + lines = _capture( + monkeypatch, + [ + _message_start(), + _content_block_start_text(), + _text_delta("hi"), + _citations_delta(cit), + _content_block_stop(), + _message_delta_end(), + _message_stop(), + ], + ) + body = _joined(lines) + payload = _citation_payload(body) + assert payload["citations"], payload + for c in payload["citations"]: + assert "_key" not in c, c + + +def test_citation_across_multiple_content_blocks_numbers_continue(monkeypatch): + """Footnote numbering is per-message, not per-content-block: + citations across separate blocks emit [1] then [2].""" + cit_a = { + "type": "char_location", + "document_index": 0, + "document_title": "Doc", + "start_char_index": 0, + "end_char_index": 5, + } + cit_b = { + "type": "char_location", + "document_index": 0, + "document_title": "Doc", + "start_char_index": 100, + "end_char_index": 105, + } + lines = _capture( + monkeypatch, + [ + _message_start(), + _content_block_start_text(), + _text_delta("first"), + _citations_delta(cit_a, index = 0), + _content_block_stop(0), + { + "type": "content_block_start", + "index": 1, + "content_block": {"type": "text", "text": ""}, + }, + _text_delta(" second", index = 1), + _citations_delta(cit_b, index = 1), + _content_block_stop(1), + _message_delta_end(), + _message_stop(), + ], + ) + body = _joined(lines) + assert "[1]" in body and "[2]" in body, body + assert body.index("[1]") < body.index("[2]") + payload = _citation_payload(body) + assert len(payload["citations"]) == 2 + + +def test_inline_marker_lands_after_text_run(monkeypatch): + """Inline ``[N]`` must land AFTER the cited text run: Anthropic + streams text then citation, so the proxy emits ``"...green.[1]"`` + not ``"[1]green"``.""" + cit = { + "type": "char_location", + "document_index": 0, + "document_title": "Doc", + "start_char_index": 0, + "end_char_index": 20, + "cited_text": "grass", + } + lines = _capture( + monkeypatch, + [ + _message_start(), + _content_block_start_text(), + _text_delta("Grass is green."), + _citations_delta(cit), + _text_delta(" Sky is blue."), + _content_block_stop(), + _message_delta_end(), + _message_stop(), + ], + ) + body = _joined(lines) + grass = body.index("Grass is green.") + marker = body.index("[1]") + sky = body.index("Sky is blue.") + assert grass < marker < sky, body + + +def test_no_synthetic_event_when_only_text_deltas(monkeypatch): + """No citations_delta means no synthetic ``document_citations`` + event; Sources panel relies on absence to suppress the section.""" + lines = _capture( + monkeypatch, + [ + _message_start(), + _content_block_start_text(), + _text_delta("Just some prose. "), + _text_delta("More prose."), + _content_block_stop(), + _message_delta_end(), + _message_stop(), + ], + ) + body = _joined(lines) + assert "document_citations" not in body + assert "[1]" not in body + + +def test_input_document_translation_enables_citations(monkeypatch): + """``input_document`` must translate to an Anthropic ``document`` + block carrying ``citations: {enabled: true}`` (both base64 and url + source branches) so upstream emits citations_delta.""" + captured_b64: dict = {} + _capture( + monkeypatch, + [ + _message_start(), + _content_block_start_text(), + _text_delta("ok"), + _content_block_stop(), + _message_delta_end(), + _message_stop(), + ], + messages = [ + { + "role": "user", + "content": [ + { + "type": "input_document", + "file_data": "data:application/pdf;base64,QUJD", + "filename": "spec.pdf", + }, + {"type": "text", "text": "summarise"}, + ], + } + ], + captured_body = captured_b64, + ) + user_msg = captured_b64["messages"][0] + doc_block = next(p for p in user_msg["content"] if p.get("type") == "document") + assert doc_block["source"]["type"] == "base64", doc_block + assert doc_block.get("citations") == {"enabled": True}, doc_block + + captured_url: dict = {} + _capture( + monkeypatch, + [ + _message_start(), + _content_block_start_text(), + _text_delta("ok"), + _content_block_stop(), + _message_delta_end(), + _message_stop(), + ], + messages = [ + { + "role": "user", + "content": [ + { + "type": "input_document", + "file_url": "https://example.com/doc.pdf", + "filename": "doc.pdf", + }, + {"type": "text", "text": "summarise"}, + ], + } + ], + captured_body = captured_url, + ) + user_msg = captured_url["messages"][0] + doc_block = next(p for p in user_msg["content"] if p.get("type") == "document") + assert doc_block["source"]["type"] == "url", doc_block + assert doc_block.get("citations") == {"enabled": True}, doc_block + + +# ── cited_text truncation + safe-url citation conversion ──────── + + +def test_cited_text_truncated_in_synthetic_event(monkeypatch): + """``cited_text`` is capped server-side so multi-KB spans do not + balloon the SSE payload.""" + from core.inference.external_provider import _CITED_TEXT_MAX_LEN + + long_quote = "x" * (_CITED_TEXT_MAX_LEN + 4000) + events = [ + { + "type": "message_start", + "message": { + "id": "msg_1", + "usage": {"input_tokens": 1, "output_tokens": 0}, + }, + }, + { + "type": "content_block_start", + "index": 0, + "content_block": {"type": "text", "text": ""}, + }, + { + "type": "content_block_delta", + "index": 0, + "delta": {"type": "text_delta", "text": "claim "}, + }, + { + "type": "content_block_delta", + "index": 0, + "delta": { + "type": "citations_delta", + "citation": { + "type": "char_location", + "document_index": 0, + "document_title": "doc", + "start_char_index": 0, + "end_char_index": 5, + "cited_text": long_quote, + }, + }, + }, + {"type": "content_block_stop", "index": 0}, + { + "type": "message_delta", + "delta": {"stop_reason": "end_turn"}, + "usage": {"output_tokens": 1}, + }, + {"type": "message_stop"}, + ] + chunks = _capture(monkeypatch, events) + tool_events = [c for c in chunks if "_toolEvent" in c and "document_citations" in c] + assert tool_events, "no document_citations tool event" + payload = json.loads(tool_events[0].split("data: ", 1)[1]) + cited = payload["_toolEvent"]["citations"][0]["cited_text"] + assert len(cited) <= _CITED_TEXT_MAX_LEN + 1, len(cited) + assert cited.endswith("…") diff --git a/studio/backend/tests/test_anthropic_fast_mode_and_refusal.py b/studio/backend/tests/test_anthropic_fast_mode_and_refusal.py new file mode 100644 index 0000000000..e7e5ec64d4 --- /dev/null +++ b/studio/backend/tests/test_anthropic_fast_mode_and_refusal.py @@ -0,0 +1,164 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0 + +"""Tests for Anthropic fast-mode wiring and streaming refusal handling. + +fast_mode=True on Opus 4.6/4.7 attaches the ``fast-mode-2026-02-01`` +beta header and sets ``speed: "fast"``; unsupported models drop both. +Streaming ``stop_reason: "refusal"`` surfaces a user notice before the +``content_filter`` finish chunk. +https://platform.claude.com/docs/en/test-and-evaluate/strengthen-guardrails/handle-streaming-refusals +""" + +import asyncio +import json + +import httpx + +from core.inference import external_provider as ep_mod +from core.inference.external_provider import ExternalProviderClient + + +def _drive(coro): + return asyncio.new_event_loop().run_until_complete(coro) + + +def _make_client() -> ExternalProviderClient: + return ExternalProviderClient( + provider_type = "anthropic", + base_url = "https://api.anthropic.com/v1", + api_key = "sk-ant-test", + ) + + +def _empty_message_sse() -> bytes: + return ( + b'event: message_start\ndata: {"type":"message_start","message":' + b'{"id":"m1","content":[],"model":"claude-opus-4-7","role":"assistant",' + b'"stop_reason":null,"usage":{"input_tokens":1,"output_tokens":1}}}\n\n' + b'event: message_delta\ndata: {"type":"message_delta",' + b'"delta":{"stop_reason":"end_turn"}}\n\n' + b'event: message_stop\ndata: {"type":"message_stop"}\n\n' + ) + + +def _refusal_sse() -> bytes: + return ( + b'event: message_start\ndata: {"type":"message_start","message":' + b'{"id":"m1","content":[],"model":"claude-opus-4-7","role":"assistant",' + b'"stop_reason":null,"usage":{"input_tokens":1,"output_tokens":1}}}\n\n' + b'event: content_block_start\ndata: {"type":"content_block_start",' + b'"index":0,"content_block":{"type":"text","text":""}}\n\n' + b'event: content_block_delta\ndata: {"type":"content_block_delta",' + b'"index":0,"delta":{"type":"text_delta","text":"Hello."}}\n\n' + b'event: content_block_stop\ndata: {"type":"content_block_stop","index":0}\n\n' + b'event: message_delta\ndata: {"type":"message_delta",' + b'"delta":{"stop_reason":"refusal"}}\n\n' + b'event: message_stop\ndata: {"type":"message_stop"}\n\n' + ) + + +def _capture(monkeypatch, sse: bytes = b"", **kwargs) -> tuple[dict, list[str]]: + """Install a MockTransport, drive one streamed call, return body+lines.""" + captured: dict = {} + + def handler(request: httpx.Request) -> httpx.Response: + captured["body"] = json.loads(request.content.decode("utf-8")) + captured["headers"] = dict(request.headers) + return httpx.Response( + 200, + content = sse or _empty_message_sse(), + headers = {"content-type": "text/event-stream"}, + ) + + monkeypatch.setattr( + ep_mod, + "_http_client", + httpx.AsyncClient(transport = httpx.MockTransport(handler)), + ) + + out_lines: list[str] = [] + + async def run(): + client = _make_client() + try: + async for line in client.stream_chat_completion( + messages = [{"role": "user", "content": "hi"}], + model = kwargs.get("model", "claude-opus-4-7"), + temperature = 0.7, + top_p = 0.95, + max_tokens = 32, + fast_mode = kwargs.get("fast_mode"), + ): + out_lines.append(line) + finally: + await client.close() + + _drive(run()) + return captured, out_lines + + +def test_fast_mode_attaches_beta_header_and_speed_on_opus_4_7(monkeypatch): + cap, _ = _capture(monkeypatch, fast_mode = True, model = "claude-opus-4-7") + assert cap["body"].get("speed") == "fast", cap["body"] + beta = cap["headers"].get("anthropic-beta", "") + assert "fast-mode-2026-02-01" in beta, beta + + +def test_fast_mode_attaches_beta_header_and_speed_on_opus_4_6(monkeypatch): + cap, _ = _capture(monkeypatch, fast_mode = True, model = "claude-opus-4-6") + assert cap["body"].get("speed") == "fast", cap["body"] + assert "fast-mode-2026-02-01" in cap["headers"].get("anthropic-beta", "") + + +def test_fast_mode_dropped_on_sonnet(monkeypatch): + cap, _ = _capture(monkeypatch, fast_mode = True, model = "claude-sonnet-4-6") + assert "speed" not in cap["body"], cap["body"] + assert "fast-mode-2026-02-01" not in cap["headers"].get("anthropic-beta", "") + + +def test_fast_mode_dropped_on_haiku(monkeypatch): + cap, _ = _capture(monkeypatch, fast_mode = True, model = "claude-haiku-4-5") + assert "speed" not in cap["body"], cap["body"] + assert "fast-mode-2026-02-01" not in cap["headers"].get("anthropic-beta", "") + + +def test_fast_mode_dropped_on_older_opus(monkeypatch): + cap, _ = _capture(monkeypatch, fast_mode = True, model = "claude-opus-4-5") + assert "speed" not in cap["body"], cap["body"] + + +def test_fast_mode_false_does_not_attach_header_or_field(monkeypatch): + cap, _ = _capture(monkeypatch, fast_mode = False) + assert "speed" not in cap["body"], cap["body"] + assert "fast-mode-2026-02-01" not in cap["headers"].get("anthropic-beta", "") + + +def test_fast_mode_none_does_not_attach_header_or_field(monkeypatch): + cap, _ = _capture(monkeypatch, fast_mode = None) + assert "speed" not in cap["body"], cap["body"] + assert "fast-mode-2026-02-01" not in cap["headers"].get("anthropic-beta", "") + + +def test_refusal_emits_user_facing_notice_and_content_filter_finish(monkeypatch): + _, lines = _capture(monkeypatch, sse = _refusal_sse()) + body = "\n".join(lines) + # User-visible refusal notice. + assert "stopped by Anthropic's safety classifier" in body, body + # OpenAI-spec finish_reason mapping. + assert '"finish_reason": "content_filter"' in body, body + # Original deltas preserved before the refusal supplement. + assert "Hello." in body, body + + +def test_refusal_emits_tool_event_for_chat_adapter_drop(monkeypatch): + """Refused turns emit an out-of-band `_toolEvent` that the chat-adapter + latches into assistant `metadata.custom.anthropicRefusal`, driving + the next-request prune. Tool event (not text) prevents spoofing. + """ + _, lines = _capture(monkeypatch, sse = _refusal_sse()) + body = "\n".join(lines) + assert '"_toolEvent": {"type": "anthropic_refusal"}' in body, body + # Visible refusal text must not embed a sentinel that could spoof + # a context reset if echoed by another assistant message. + assert "studio:anthropic-refusal" not in body, body diff --git a/studio/backend/tests/test_anthropic_fast_mode_edge.py b/studio/backend/tests/test_anthropic_fast_mode_edge.py new file mode 100644 index 0000000000..0052cb94ad --- /dev/null +++ b/studio/backend/tests/test_anthropic_fast_mode_edge.py @@ -0,0 +1,442 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0 + +"""Edge-case coverage for the Anthropic fast-mode + refusal wiring. + +Complements ``test_anthropic_fast_mode_and_refusal.py`` (happy path) +with dated snapshots, strict opt-in (future Opus families do not +auto-enable), multi-beta header merging, refusal stream ordering, and +the non-destruction guarantee for unset/None fast_mode. +""" + +import asyncio +import json +import re + +import httpx + +from core.inference import external_provider as ep_mod +from core.inference.external_provider import ExternalProviderClient + + +def _drive(coro): + return asyncio.new_event_loop().run_until_complete(coro) + + +def _make_client() -> ExternalProviderClient: + return ExternalProviderClient( + provider_type = "anthropic", + base_url = "https://api.anthropic.com/v1", + api_key = "sk-ant-test", + ) + + +def _empty_message_sse(model: str = "claude-opus-4-7") -> bytes: + return ( + b'event: message_start\ndata: {"type":"message_start","message":' + b'{"id":"m1","content":[],"model":"' + model.encode() + b'",' + b'"role":"assistant","stop_reason":null,"usage":' + b'{"input_tokens":1,"output_tokens":1}}}\n\n' + b'event: message_delta\ndata: {"type":"message_delta",' + b'"delta":{"stop_reason":"end_turn"}}\n\n' + b'event: message_stop\ndata: {"type":"message_stop"}\n\n' + ) + + +def _refusal_sse(model: str = "claude-opus-4-7") -> bytes: + return ( + b'event: message_start\ndata: {"type":"message_start","message":' + b'{"id":"m1","content":[],"model":"' + model.encode() + b'",' + b'"role":"assistant","stop_reason":null,"usage":' + b'{"input_tokens":1,"output_tokens":1}}}\n\n' + b'event: content_block_start\ndata: {"type":"content_block_start",' + b'"index":0,"content_block":{"type":"text","text":""}}\n\n' + b'event: content_block_delta\ndata: {"type":"content_block_delta",' + b'"index":0,"delta":{"type":"text_delta","text":"Hello."}}\n\n' + b'event: content_block_stop\ndata: {"type":"content_block_stop","index":0}\n\n' + b'event: message_delta\ndata: {"type":"message_delta",' + b'"delta":{"stop_reason":"refusal"}}\n\n' + b'event: message_stop\ndata: {"type":"message_stop"}\n\n' + ) + + +def _capture(monkeypatch, sse: bytes = b"", **kwargs) -> tuple[dict, list[str]]: + """Install a MockTransport, drive one streamed call, return body+lines.""" + captured: dict = {} + + def handler(request: httpx.Request) -> httpx.Response: + captured["body"] = json.loads(request.content.decode("utf-8")) + captured["headers"] = dict(request.headers) + return httpx.Response( + 200, + content = sse or _empty_message_sse(kwargs.get("model", "claude-opus-4-7")), + headers = {"content-type": "text/event-stream"}, + ) + + monkeypatch.setattr( + ep_mod, + "_http_client", + httpx.AsyncClient(transport = httpx.MockTransport(handler)), + ) + + out_lines: list[str] = [] + + async def run(): + client = _make_client() + try: + extra = {} + for key in ( + "enabled_tools", + "compaction_threshold", + "fast_mode", + ): + if key in kwargs: + extra[key] = kwargs[key] + async for line in client.stream_chat_completion( + messages = [{"role": "user", "content": "hi"}], + model = kwargs.get("model", "claude-opus-4-7"), + temperature = 0.7, + top_p = 0.95, + max_tokens = 32, + **extra, + ): + out_lines.append(line) + finally: + await client.close() + + _drive(run()) + return captured, out_lines + + +# ──────────────────────────── dated snapshot prefix ──────────────────────────── +def test_fast_mode_attaches_on_dated_opus_4_7_snapshot(monkeypatch): + """Dated snapshot ``claude-opus-4-7-2026-02-01`` must match the prefix.""" + cap, _ = _capture(monkeypatch, fast_mode = True, model = "claude-opus-4-7-2026-02-01") + assert cap["body"].get("speed") == "fast", cap["body"] + assert "fast-mode-2026-02-01" in cap["headers"].get("anthropic-beta", "") + + +def test_fast_mode_attaches_on_dated_opus_4_6_snapshot(monkeypatch): + cap, _ = _capture(monkeypatch, fast_mode = True, model = "claude-opus-4-6-2026-02-01") + assert cap["body"].get("speed") == "fast", cap["body"] + assert "fast-mode-2026-02-01" in cap["headers"].get("anthropic-beta", "") + + +# ──────────────────────────── strict opt-in semantics ──────────────────────────── +def test_fast_mode_does_not_auto_enable_on_future_opus_4_8(monkeypatch): + """Future ``claude-opus-4-8`` must not auto-enable; opt-in per family.""" + cap, _ = _capture(monkeypatch, fast_mode = True, model = "claude-opus-4-8") + assert "speed" not in cap["body"], cap["body"] + assert "fast-mode-2026-02-01" not in cap["headers"].get("anthropic-beta", "") + + +def test_fast_mode_does_not_auto_enable_on_future_opus_5(monkeypatch): + cap, _ = _capture(monkeypatch, fast_mode = True, model = "claude-opus-5") + assert "speed" not in cap["body"], cap["body"] + assert "fast-mode-2026-02-01" not in cap["headers"].get("anthropic-beta", "") + + +def test_fast_mode_does_not_auto_enable_on_sonnet_dated_snapshot(monkeypatch): + """Sonnet snapshots share the compaction prefix but not fast_mode.""" + cap, _ = _capture(monkeypatch, fast_mode = True, model = "claude-sonnet-4-6-2026-02-01") + assert "speed" not in cap["body"], cap["body"] + assert "fast-mode-2026-02-01" not in cap["headers"].get("anthropic-beta", "") + + +# ──────────────────────────── beta header merge ──────────────────────────── +def _beta_parts(headers: dict) -> list[str]: + raw = headers.get("anthropic-beta", "") + return [p.strip() for p in raw.split(",") if p.strip()] + + +def test_fast_mode_merges_with_code_execution_beta(monkeypatch): + """fast_mode + code_execution -> two comma-separated betas, no overwrite.""" + cap, _ = _capture( + monkeypatch, + fast_mode = True, + model = "claude-opus-4-7", + enabled_tools = ["code_execution"], + ) + parts = _beta_parts(cap["headers"]) + assert "fast-mode-2026-02-01" in parts, cap["headers"] + assert any(p.startswith("code-execution-") for p in parts), cap["headers"] + # No duplicates. + assert len(parts) == len(set(parts)), parts + + +def test_fast_mode_merges_with_compaction_beta(monkeypatch): + """fast_mode + compaction_threshold >= 50K -> both betas present.""" + cap, _ = _capture( + monkeypatch, + fast_mode = True, + model = "claude-opus-4-7", + compaction_threshold = 100_000, + ) + parts = _beta_parts(cap["headers"]) + assert "fast-mode-2026-02-01" in parts, cap["headers"] + assert "compact-2026-01-12" in parts, cap["headers"] + + +def test_fast_mode_merges_with_code_execution_and_compaction(monkeypatch): + """Three betas coexist in one comma-separated header, no duplicates.""" + cap, _ = _capture( + monkeypatch, + fast_mode = True, + model = "claude-opus-4-7", + enabled_tools = ["code_execution"], + compaction_threshold = 100_000, + ) + parts = _beta_parts(cap["headers"]) + assert "fast-mode-2026-02-01" in parts + assert "compact-2026-01-12" in parts + assert any(p.startswith("code-execution-") for p in parts), parts + assert len(parts) >= 3 + assert len(parts) == len(set(parts)), parts + + +def test_fast_mode_beta_value_is_pinned(monkeypatch): + """Pin the exact beta tag ``fast-mode-2026-02-01`` from the docs.""" + cap, _ = _capture(monkeypatch, fast_mode = True, model = "claude-opus-4-7") + parts = _beta_parts(cap["headers"]) + assert "fast-mode-2026-02-01" in parts, parts + # Reject obvious typos. + assert not any(p.startswith("fastmode-") for p in parts), parts + assert not any("fast_mode" in p for p in parts), parts + + +# ──────────────────────────── non-destruction guarantee ──────────────────────────── +def test_fast_mode_unset_is_byte_identical_to_omitted(monkeypatch): + """``fast_mode=None`` must produce the same body/headers as omission.""" + cap_none, _ = _capture(monkeypatch, fast_mode = None, model = "claude-opus-4-7") + + # Re-run without passing fast_mode at all. + captured: dict = {} + + def handler(request: httpx.Request) -> httpx.Response: + captured["body"] = json.loads(request.content.decode("utf-8")) + captured["headers"] = dict(request.headers) + return httpx.Response( + 200, + content = _empty_message_sse(), + headers = {"content-type": "text/event-stream"}, + ) + + monkeypatch.setattr( + ep_mod, + "_http_client", + httpx.AsyncClient(transport = httpx.MockTransport(handler)), + ) + + async def run(): + client = _make_client() + try: + async for _ in client.stream_chat_completion( + messages = [{"role": "user", "content": "hi"}], + model = "claude-opus-4-7", + temperature = 0.7, + top_p = 0.95, + max_tokens = 32, + ): + pass + finally: + await client.close() + + _drive(run()) + + assert cap_none["body"] == captured["body"], (cap_none["body"], captured["body"]) + # Headers can vary by httpx-injected fields (host, connection); compare + # the load-bearing ones. + for key in ("anthropic-version", "x-api-key", "content-type"): + assert cap_none["headers"].get(key) == captured["headers"].get(key), key + assert "anthropic-beta" not in cap_none["headers"] + assert "anthropic-beta" not in captured["headers"] + assert "speed" not in cap_none["body"] + assert "speed" not in captured["body"] + + +def test_fast_mode_false_on_opus_4_7_byte_identical_to_unset(monkeypatch): + """``fast_mode=False`` produces the same outbound shape as unset.""" + cap_false, _ = _capture(monkeypatch, fast_mode = False, model = "claude-opus-4-7") + assert "speed" not in cap_false["body"], cap_false["body"] + assert "fast-mode-2026-02-01" not in cap_false["headers"].get("anthropic-beta", "") + + +# ──────────────────────────── refusal stream ordering ──────────────────────────── +def test_refusal_notice_appears_before_content_filter_chunk(monkeypatch): + """The notice content delta must precede the finish_reason chunk.""" + _, lines = _capture(monkeypatch, sse = _refusal_sse(), model = "claude-opus-4-7") + notice_idx = next(i for i, l in enumerate(lines) if "stopped by Anthropic" in l) + filter_idx = next( + i for i, l in enumerate(lines) if '"finish_reason": "content_filter"' in l + ) + assert notice_idx < filter_idx, (notice_idx, filter_idx, lines) + + +def test_refusal_tool_event_emitted_exactly_once(monkeypatch): + """A single refusal emits the chat-adapter drop signal exactly once.""" + _, lines = _capture(monkeypatch, sse = _refusal_sse()) + body = "\n".join(lines) + count = body.count('"_toolEvent": {"type": "anthropic_refusal"}') + assert count == 1, (count, body) + + +def test_refusal_text_carries_no_html_sentinel(monkeypatch): + """Visible refusal text must not embed a ``studio:anthropic-refusal`` + sentinel; the drop signal rides _toolEvent only.""" + _, lines = _capture(monkeypatch, sse = _refusal_sse()) + body = "\n".join(lines) + assert "studio:anthropic-refusal" not in body, body + + +def test_refusal_handling_works_on_sonnet_model(monkeypatch): + """Refusal handling is provider-side; Sonnet refusals must also surface.""" + _, lines = _capture( + monkeypatch, sse = _refusal_sse("claude-sonnet-4-6"), model = "claude-sonnet-4-6" + ) + body = "\n".join(lines) + assert "stopped by Anthropic's safety classifier" in body, body + assert '"_toolEvent": {"type": "anthropic_refusal"}' in body, body + assert '"finish_reason": "content_filter"' in body, body + + +def test_refusal_preserves_partial_assistant_text(monkeypatch): + """Partial deltas already streamed must precede the refusal notice.""" + _, lines = _capture(monkeypatch, sse = _refusal_sse(), model = "claude-opus-4-7") + body = "\n".join(lines) + hello_idx = body.index("Hello.") + notice_idx = body.index("stopped by Anthropic") + assert hello_idx < notice_idx, (hello_idx, notice_idx) + + +def test_refusal_chunk_is_proper_openai_delta_shape(monkeypatch): + """The notice rides ``choices[0].delta.content`` (not a finish chunk); + OpenAI-spec clients treat it as ordinary streamed text.""" + _, lines = _capture(monkeypatch, sse = _refusal_sse(), model = "claude-opus-4-7") + # Find the chunk that carries the refusal text. + notice_chunk = None + for line in lines: + if line.startswith("data: ") and "stopped by Anthropic" in line: + notice_chunk = json.loads(line[len("data: ") :]) + break + assert notice_chunk is not None, lines + choice = notice_chunk["choices"][0] + assert "delta" in choice and "content" in choice["delta"], notice_chunk + # Must NOT carry a finish_reason itself -- that comes on the next + # chunk. + assert choice.get("finish_reason") in (None,), notice_chunk + # Refusal text is plain-spoken; no embedded sentinel. + assert "studio:anthropic-refusal" not in choice["delta"]["content"] + + +def test_refusal_tool_event_chunk_shape(monkeypatch): + """Drop signal rides a Studio `_toolEvent` envelope (delta={}, + finish_reason=null); the frontend latches on + `_toolEvent.type == "anthropic_refusal"`.""" + _, lines = _capture(monkeypatch, sse = _refusal_sse(), model = "claude-opus-4-7") + refusal_chunk = None + for line in lines: + if line.startswith("data: ") and "anthropic_refusal" in line: + refusal_chunk = json.loads(line[len("data: ") :]) + break + assert refusal_chunk is not None, lines + assert refusal_chunk["_toolEvent"] == {"type": "anthropic_refusal"}, refusal_chunk + choice = refusal_chunk["choices"][0] + assert choice["delta"] == {}, refusal_chunk + assert choice["finish_reason"] is None, refusal_chunk + + +# ──────────────────────────── future-proofing ──────────────────────────── +def test_fast_mode_prefix_tuple_matches_capability_doc(monkeypatch): + """Tuple must exactly match the two families in the upstream docs: + https://platform.claude.com/docs/en/build-with-claude/fast-mode.""" + from core.inference.external_provider import _ANTHROPIC_FAST_MODE_PREFIXES + + assert set(_ANTHROPIC_FAST_MODE_PREFIXES) == { + "claude-opus-4-7", + "claude-opus-4-6", + }, _ANTHROPIC_FAST_MODE_PREFIXES + + +def test_fast_mode_speed_field_value_is_literal_fast(monkeypatch): + """Pin the wire value to the literal string ``"fast"``.""" + cap, _ = _capture(monkeypatch, fast_mode = True, model = "claude-opus-4-7") + assert cap["body"]["speed"] == "fast", cap["body"] + + +def test_fast_mode_dropped_on_opus_4_5_dated_snapshot(monkeypatch): + """Previous-family snapshots like ``claude-opus-4-5-2025-...`` must not match.""" + cap, _ = _capture(monkeypatch, fast_mode = True, model = "claude-opus-4-5-2025-08-01") + assert "speed" not in cap["body"], cap["body"] + assert "fast-mode-2026-02-01" not in cap["headers"].get("anthropic-beta", "") + + +def test_fast_mode_rejects_prefix_collision_4_70(monkeypatch): + """IDs like ``claude-opus-4-70`` / ``-4-7b`` must not match the prefix.""" + cap, _ = _capture(monkeypatch, fast_mode = True, model = "claude-opus-4-70") + assert "speed" not in cap["body"], cap["body"] + assert "fast-mode-2026-02-01" not in cap["headers"].get("anthropic-beta", "") + + +def test_fast_mode_rejects_prefix_collision_4_7b(monkeypatch): + cap, _ = _capture(monkeypatch, fast_mode = True, model = "claude-opus-4-7b") + assert "speed" not in cap["body"], cap["body"] + assert "fast-mode-2026-02-01" not in cap["headers"].get("anthropic-beta", "") + + +def test_fast_mode_rejects_prefix_collision_4_6_extra(monkeypatch): + cap, _ = _capture(monkeypatch, fast_mode = True, model = "claude-opus-4-60") + assert "speed" not in cap["body"], cap["body"] + assert "fast-mode-2026-02-01" not in cap["headers"].get("anthropic-beta", "") + + +# ──────────────────────────── usage.speed propagation ──────────────────────────── +def _fast_speed_sse(model: str = "claude-opus-4-7", speed: str = "fast") -> bytes: + return ( + b'event: message_start\ndata: {"type":"message_start","message":' + b'{"id":"m1","content":[],"model":"' + model.encode() + b'",' + b'"role":"assistant","stop_reason":null,"usage":' + b'{"input_tokens":4,"output_tokens":1}}}\n\n' + b'event: content_block_start\ndata: {"type":"content_block_start",' + b'"index":0,"content_block":{"type":"text","text":""}}\n\n' + b'event: content_block_delta\ndata: {"type":"content_block_delta",' + b'"index":0,"delta":{"type":"text_delta","text":"hi"}}\n\n' + b'event: content_block_stop\ndata: {"type":"content_block_stop","index":0}\n\n' + b'event: message_delta\ndata: {"type":"message_delta",' + b'"delta":{"stop_reason":"end_turn"},' + b'"usage":{"output_tokens":5,"speed":"' + speed.encode() + b'"}}\n\n' + b'event: message_stop\ndata: {"type":"message_stop"}\n\n' + ) + + +def test_usage_speed_propagates_to_final_usage_chunk_fast(monkeypatch): + """``usage.speed == "fast"`` from upstream must reach the Studio usage chunk.""" + _, lines = _capture(monkeypatch, sse = _fast_speed_sse(speed = "fast")) + usage_lines = [l for l in lines if l.startswith("data: ") and '"usage"' in l] + assert usage_lines, lines + parsed = [json.loads(l[len("data: ") :]) for l in usage_lines] + speeds = [p["usage"].get("speed") for p in parsed if "usage" in p] + assert "fast" in speeds, parsed + + +def test_usage_speed_propagates_to_final_usage_chunk_standard(monkeypatch): + _, lines = _capture(monkeypatch, sse = _fast_speed_sse(speed = "standard")) + parsed = [ + json.loads(l[len("data: ") :]) + for l in lines + if l.startswith("data: ") and '"usage"' in l + ] + speeds = [p["usage"].get("speed") for p in parsed if "usage" in p] + assert "standard" in speeds, parsed + + +def test_usage_speed_absent_when_anthropic_does_not_report(monkeypatch): + """Studio must not invent ``usage.speed`` when upstream omits it.""" + _, lines = _capture(monkeypatch) + parsed = [ + json.loads(l[len("data: ") :]) + for l in lines + if l.startswith("data: ") and '"usage"' in l + ] + for p in parsed: + usage = p.get("usage") or {} + assert "speed" not in usage, p diff --git a/studio/backend/tests/test_anthropic_web_fetch.py b/studio/backend/tests/test_anthropic_web_fetch.py index cdb5f6254c..da10d679eb 100644 --- a/studio/backend/tests/test_anthropic_web_fetch.py +++ b/studio/backend/tests/test_anthropic_web_fetch.py @@ -2,26 +2,12 @@ # Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0 """ -Unit tests for Anthropic's server-side `web_fetch_20250910` tool -translation in `_stream_anthropic`. - -Covers: -- Request body: when ``enabled_tools=["web_fetch"]``, the outbound - ``tools`` array carries ``{"type":"web_fetch_20250910", - "name":"web_fetch", "max_uses":5}``. No beta header is required. -- Combined request: ``enabled_tools=["web_search","web_fetch", - "code_execution"]`` sends all three tool entries. -- Disabled by default: with ``enabled_tools=["web_search"]`` (or None), - the body does NOT carry a web_fetch entry. -- SSE translation (success): a `web_fetch` server_tool_use streaming - ``{"url": "..."}`` followed by a `web_fetch_tool_result` block with - a document source emits one ``tool_start`` and one ``tool_end`` - `_toolEvent`. The ``tool_start.arguments.url`` matches the fetched - URL and the ``tool_end.result`` carries the Title / URL / snippet - prefix the source-pill renderer expects. -- SSE translation (error): a `web_fetch_tool_error` with - ``error_code="url_not_accessible"`` renders as ``"Error: - url_not_accessible"`` in the tool_end result. +Unit tests for Anthropic's `web_fetch_20250910` / `web_fetch_20260209` +translation in ``_stream_anthropic``. Covers request body emission +(version picked by ``_anthropic_web_fetch_version``: ``_20260209`` for +Opus 4.6/4.7 + Sonnet 4.6, ``_20250910`` otherwise), combined tool +requests, off-by-default behavior, and SSE translation of success and +``url_not_accessible`` error paths into ``tool_start`` / ``tool_end``. """ import asyncio @@ -117,8 +103,9 @@ def test_web_fetch_tool_appended_to_request_body(monkeypatch): body = captured["body"] tools = body.get("tools") or [] + # claude-opus-4-7 routes web_fetch to _20260209 (dynamic filtering). assert { - "type": "web_fetch_20250910", + "type": "web_fetch_20260209", "name": "web_fetch", "max_uses": 5, } in tools @@ -157,13 +144,10 @@ def test_web_fetch_combined_with_web_search_and_code_execution(monkeypatch): tools = captured["body"].get("tools") or [] tool_types = [t.get("type") for t in tools] - # After PR 5679's per-model tool version dispatch landed, - # claude-opus-4-7 routes web_search to the _20260209 variant and - # code_execution to _20260120. web_fetch still hardcodes - # _20250910 today; see follow-up to thread it through - # _anthropic_web_fetch_version. + # claude-opus-4-7 routes web_search and web_fetch to _20260209 + # and code_execution to _20260120 (per PR 5679 dispatch). assert "web_search_20260209" in tool_types, tool_types - assert "web_fetch_20250910" in tool_types, tool_types + assert "web_fetch_20260209" in tool_types, tool_types assert "code_execution_20260120" in tool_types, tool_types # Code-execution still adds its beta flag; web_fetch must not # have accidentally stripped it. @@ -199,7 +183,9 @@ def test_no_web_fetch_tool_when_pill_off(monkeypatch): _drive(run()) tools = captured["body"].get("tools") or [] - assert all(t.get("type") != "web_fetch_20250910" for t in tools) + assert all( + t.get("type") not in ("web_fetch_20250910", "web_fetch_20260209") for t in tools + ) # ── SSE translation ───────────────────────────────────────────────── @@ -365,7 +351,9 @@ def test_web_fetch_error_renders_error_code(monkeypatch): def _finish_reasons(lines: list[str]) -> list: - """Return the finish_reason fields from every chat.completion.chunk.""" + """Return non-null finish_reason fields from each chat.completion.chunk. + Mid-stream content deltas carry ``finish_reason: None`` and are skipped + (the refusal path emits a notice delta before the content_filter chunk).""" out: list = [] for line in lines: if not line.startswith("data:"): @@ -380,8 +368,9 @@ def _finish_reasons(lines: list[str]) -> list: if parsed.get("object") != "chat.completion.chunk": continue for choice in parsed.get("choices") or []: - if "finish_reason" in choice: - out.append(choice["finish_reason"]) + reason = choice.get("finish_reason") + if reason is not None: + out.append(reason) return out diff --git a/studio/backend/tests/test_index_bootstrap_origin.py b/studio/backend/tests/test_index_bootstrap_origin.py new file mode 100644 index 0000000000..89f7613ee4 --- /dev/null +++ b/studio/backend/tests/test_index_bootstrap_origin.py @@ -0,0 +1,144 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0 + +"""Regression coverage for the bootstrap-pw cross-origin leak (PR 5739). +``_is_same_origin_request`` gates ``_inject_bootstrap`` so the seeded +admin password only ships to same-origin callers. +""" + +import os +from types import SimpleNamespace +from unittest.mock import MagicMock + +import pytest + + +def _build_request(host: str, origin: str | None, scheme: str = "http") -> MagicMock: + request = MagicMock() + request.url.scheme = scheme + request.url.netloc = host + request.headers = {"origin": origin} if origin is not None else {} + return request + + +def test_is_same_origin_request_missing_origin_is_same_origin(monkeypatch): + from main import _is_same_origin_request + + req = _build_request("127.0.0.1:8888", origin = None) + assert _is_same_origin_request(req) is True + + +def test_is_same_origin_request_matching_origin_is_same_origin(): + from main import _is_same_origin_request + + req = _build_request("127.0.0.1:8888", origin = "http://127.0.0.1:8888") + assert _is_same_origin_request(req) is True + + +def test_is_same_origin_request_evil_origin_is_cross_origin(): + from main import _is_same_origin_request + + req = _build_request("127.0.0.1:8888", origin = "https://evil.example") + assert _is_same_origin_request(req) is False + + +def test_is_same_origin_request_scheme_mismatch_is_cross_origin(): + # https origin against an http listener is not same-origin. + from main import _is_same_origin_request + + req = _build_request("127.0.0.1:8888", origin = "https://127.0.0.1:8888") + assert _is_same_origin_request(req) is False + + +def test_is_same_origin_request_port_mismatch_is_cross_origin(): + # Same host different port is not same-origin per the web platform. + from main import _is_same_origin_request + + req = _build_request("127.0.0.1:8888", origin = "http://127.0.0.1:5173") + assert _is_same_origin_request(req) is False + + +# ── Canonicalisation: default-port stripping + case folding ───────── + + +def test_is_same_origin_request_https_default_port_stripped_on_origin(): + """RFC 6454 strips default ports on Origin; Starlette's netloc may still + carry ``:443``. Canonicalise both sides so this stays same-origin. + """ + from main import _is_same_origin_request + + req = _build_request( + "example.com:443", origin = "https://example.com", scheme = "https" + ) + assert _is_same_origin_request(req) is True + + +def test_is_same_origin_request_http_default_port_stripped_on_origin(): + from main import _is_same_origin_request + + req = _build_request("example.com:80", origin = "http://example.com") + assert _is_same_origin_request(req) is True + + +def test_is_same_origin_request_default_port_present_on_origin(): + """Mirror case: Origin carries the default port, netloc doesn't. Same-origin.""" + from main import _is_same_origin_request + + req = _build_request( + "example.com", origin = "https://example.com:443", scheme = "https" + ) + assert _is_same_origin_request(req) is True + + +def test_is_same_origin_request_host_case_insensitive(): + """Host portion is case-insensitive per RFC 3986.""" + from main import _is_same_origin_request + + req = _build_request("example.com", origin = "http://EXAMPLE.com") + assert _is_same_origin_request(req) is True + + +def test_is_same_origin_request_scheme_case_insensitive(): + """Scheme portion is case-insensitive per RFC 3986.""" + from main import _is_same_origin_request + + req = _build_request("example.com", origin = "HTTP://example.com") + assert _is_same_origin_request(req) is True + + +def test_is_same_origin_request_null_origin_is_cross_origin(): + """Sandboxed iframes / file:// pages send ``Origin: null``; cross-origin.""" + from main import _is_same_origin_request + + req = _build_request("example.com", origin = "null") + assert _is_same_origin_request(req) is False + + +def test_is_same_origin_request_unparseable_origin_is_cross_origin(): + """Garbage values without a host fall to cross-origin; a malformed header + must not leak the bootstrap. + """ + from main import _is_same_origin_request + + req = _build_request("example.com", origin = "not-a-url") + assert _is_same_origin_request(req) is False + + +def test_is_same_origin_request_userinfo_in_netloc_ignored(): + """``user:pass@host:port`` netlocs (RFC 3986) must compare equal to the + credentials-less Origin. + """ + from main import _is_same_origin_request + + req = _build_request("user:pass@example.com:80", origin = "http://example.com") + assert _is_same_origin_request(req) is True + + +def test_is_same_origin_request_explicit_non_default_port_still_mismatch(): + """Canonicalisation does NOT collapse non-default ports to default.""" + from main import _is_same_origin_request + + req = _build_request( + "example.com", origin = "https://example.com:9999", scheme = "https" + ) + assert _is_same_origin_request(req) is False diff --git a/studio/backend/tests/test_index_bootstrap_origin_extra.py b/studio/backend/tests/test_index_bootstrap_origin_extra.py new file mode 100644 index 0000000000..aea6b36a96 --- /dev/null +++ b/studio/backend/tests/test_index_bootstrap_origin_extra.py @@ -0,0 +1,196 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0 + +"""Extra edge-case coverage for the bootstrap-pw cross-origin gate. +Companion to ``test_index_bootstrap_origin.py``: IPv6 netlocs, opaque +origins (``data:``, ``blob:``), comma-joined multi-Origin headers, and +the ``localhost`` vs ``127.0.0.1`` distinct-origin rule. +""" + +from unittest.mock import MagicMock + + +def _build_request(host: str, origin, scheme: str = "http") -> MagicMock: + request = MagicMock() + request.url.scheme = scheme + request.url.netloc = host + request.headers = {"origin": origin} if origin is not None else {} + return request + + +# ── IPv6 ──────────────────────────────────────────────────────────── + + +def test_is_same_origin_request_ipv6_loopback_same_origin(): + """Studio supports ``-H ::1`` binds; netloc is ``[::1]:8902``. Bare + ``partition(":")`` mis-parses the bracketed form and would refuse the + bootstrap on legitimate same-origin nav. + """ + from main import _is_same_origin_request + + req = _build_request("[::1]:8902", origin = "http://[::1]:8902") + assert _is_same_origin_request(req) is True + + +def test_is_same_origin_request_ipv6_full_address_same_origin(): + from main import _is_same_origin_request + + req = _build_request( + "[2001:db8::1]:8443", + origin = "https://[2001:db8::1]:8443", + scheme = "https", + ) + assert _is_same_origin_request(req) is True + + +def test_is_same_origin_request_ipv6_default_port_stripped(): + """Browser drops :80 on ``http://[::1]``.""" + from main import _is_same_origin_request + + req = _build_request("[::1]:80", origin = "http://[::1]") + assert _is_same_origin_request(req) is True + + +def test_is_same_origin_request_ipv6_case_insensitive(): + """Hex digits in IPv6 are case-insensitive per RFC 5952.""" + from main import _is_same_origin_request + + req = _build_request( + "[2001:DB8::1]:8443", + origin = "https://[2001:db8::1]:8443", + scheme = "https", + ) + assert _is_same_origin_request(req) is True + + +def test_is_same_origin_request_ipv6_different_host_cross_origin(): + from main import _is_same_origin_request + + req = _build_request("[::1]:8902", origin = "http://[2001:db8::1]:8902") + assert _is_same_origin_request(req) is False + + +def test_is_same_origin_request_ipv6_port_mismatch_cross_origin(): + from main import _is_same_origin_request + + req = _build_request("[::1]:8902", origin = "http://[::1]:9999") + assert _is_same_origin_request(req) is False + + +def test_is_same_origin_request_ipv6_userinfo_stripped(): + from main import _is_same_origin_request + + req = _build_request("user:pass@[::1]:8902", origin = "http://[::1]:8902") + assert _is_same_origin_request(req) is True + + +# ── Opaque origins (data:, blob:) ─────────────────────────────────── + + +def test_is_same_origin_request_data_url_origin_is_cross_origin(): + """``data:`` URLs are opaque origins (HTML living standard); no host, + never same-origin. + """ + from main import _is_same_origin_request + + req = _build_request( + "127.0.0.1:8902", origin = "data:text/html," + ) + assert _is_same_origin_request(req) is False + + +def test_is_same_origin_request_blob_url_origin_is_cross_origin(): + """``blob:`` URLs carry the inner origin only in non-canonical form; the + canonical comparison rejects them. + """ + from main import _is_same_origin_request + + req = _build_request("127.0.0.1:8902", origin = "blob:http://127.0.0.1:8902/uuid") + assert _is_same_origin_request(req) is False + + +def test_is_same_origin_request_file_url_origin_is_cross_origin(): + """``file://`` pages usually send ``Origin: null``; historical engines + sent ``Origin: file://``. Neither is same-origin vs an http listener. + """ + from main import _is_same_origin_request + + req = _build_request("127.0.0.1:8902", origin = "file://") + assert _is_same_origin_request(req) is False + + +# ── Multi-Origin header (comma-joined by Starlette) ──────────────── + + +def test_is_same_origin_request_comma_joined_origins_cross_origin(): + """Starlette concatenates repeated headers with ``, ``; the canonical + parser can't safely split this, so it falls to cross-origin. + """ + from main import _is_same_origin_request + + req = _build_request( + "127.0.0.1:8902", + origin = "http://127.0.0.1:8902, http://evil.example", + ) + assert _is_same_origin_request(req) is False + + +# ── localhost vs 127.0.0.1 (distinct origins per web platform) ────── + + +def test_is_same_origin_request_localhost_vs_127_is_cross_origin(): + """Browsers treat ``localhost`` and ``127.0.0.1`` as distinct origins; + the canonical comparison must not DNS-collapse them. + """ + from main import _is_same_origin_request + + req = _build_request("127.0.0.1:8902", origin = "http://localhost:8902") + assert _is_same_origin_request(req) is False + + +def test_is_same_origin_request_127_vs_localhost_is_cross_origin(): + from main import _is_same_origin_request + + req = _build_request("localhost:8902", origin = "http://127.0.0.1:8902") + assert _is_same_origin_request(req) is False + + +# ── urlparse ValueError robustness ───────────────────────────────── + + +def test_is_same_origin_request_malformed_ipv6_bracket_is_cross_origin(): + """``urlparse`` raises ``ValueError('Invalid IPv6 URL')`` on unclosed + brackets (CVE-2024-11168 hardening). The gate must swallow and fall to + cross-origin rather than 500 the SPA handler. + """ + from main import _is_same_origin_request + + req = _build_request("127.0.0.1:8902", origin = "http://[malformed") + assert _is_same_origin_request(req) is False + + +def test_is_same_origin_request_invalid_ipv6_address_is_cross_origin(): + """Bracketed but invalid IPv6 (e.g. ``[::g]``) also raises + ``ValueError`` inside ``urlparse``.""" + from main import _is_same_origin_request + + req = _build_request("127.0.0.1:8902", origin = "http://[::g]:8902") + assert _is_same_origin_request(req) is False + + +def test_is_same_origin_request_bracket_with_trailing_garbage_is_cross_origin(): + """Text after the closing bracket also raises inside ``urlparse``.""" + from main import _is_same_origin_request + + req = _build_request("127.0.0.1:8902", origin = "http://[2001:db8::1]extra:8902") + assert _is_same_origin_request(req) is False + + +def test_is_same_origin_request_empty_origin_header_is_cross_origin(): + """Explicit empty ``Origin:`` is not a valid serialised origin and must + not be conflated with a missing header; cross-origin, bootstrap withheld. + """ + from main import _is_same_origin_request + + req = _build_request("127.0.0.1:8902", origin = "") + assert _is_same_origin_request(req) is False diff --git a/studio/backend/tests/test_multimodal_document.py b/studio/backend/tests/test_multimodal_document.py index a431b78352..4d7528d238 100644 --- a/studio/backend/tests/test_multimodal_document.py +++ b/studio/backend/tests/test_multimodal_document.py @@ -117,6 +117,8 @@ def test_anthropic_base64_pdf_becomes_document_block(monkeypatch): types = [p.get("type") for p in parts] assert "document" in types, parts doc = _strip_cache(next(p for p in parts if p.get("type") == "document")) + # citations: {enabled: true} opts into Anthropic's natural-citation + # pipeline; without it the citations_delta handler is a no-op. assert doc == { "type": "document", "source": { @@ -124,6 +126,7 @@ def test_anthropic_base64_pdf_becomes_document_block(monkeypatch): "media_type": "application/pdf", "data": _TINY_PDF_B64, }, + "citations": {"enabled": True}, "title": "paper.pdf", } @@ -151,6 +154,7 @@ def test_anthropic_url_pdf_becomes_document_block(monkeypatch): assert doc == { "type": "document", "source": {"type": "url", "url": "https://example.com/doc.pdf"}, + "citations": {"enabled": True}, } @@ -255,6 +259,7 @@ def test_anthropic_empty_data_uri_falls_back_to_file_url(monkeypatch): assert doc == { "type": "document", "source": {"type": "url", "url": "https://example.com/doc.pdf"}, + "citations": {"enabled": True}, "title": "doc.pdf", } @@ -283,6 +288,7 @@ def test_anthropic_whitespace_only_data_uri_falls_back_to_file_url(monkeypatch): assert doc == { "type": "document", "source": {"type": "url", "url": "https://example.com/doc.pdf"}, + "citations": {"enabled": True}, } diff --git a/studio/backend/tests/test_openai_citation_markers.py b/studio/backend/tests/test_openai_citation_markers.py new file mode 100644 index 0000000000..ccc17be329 --- /dev/null +++ b/studio/backend/tests/test_openai_citation_markers.py @@ -0,0 +1,251 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0 + +"""Tests for the OpenAI Responses-API citation marker rewriter. + +The stream interleaves text deltas with ``\\ue200cite\\ue202SOURCE_ID\\ue201`` +markers. The rewriter resolves each to `[N](URL)` when the annotation has +arrived and drops it otherwise; the URL list still flows to Sources via +`_record_url_citation`. + +Reference: https://developers.openai.com/api/docs/guides/citation-formatting +""" + +import pytest + +from core.inference.external_provider import ( + _replace_openai_citation_markers, + _rewrite_citation_markers_partial, +) + + +# Citation marker control codepoints (private-use area): +CITE_START = "" +CITE_STOP = "" +CITE_DELIM = "" + + +def _marker(source_id: str, locator: str | None = None) -> str: + payload = f"{CITE_START}cite{CITE_DELIM}{source_id}" + if locator: + payload = f"{payload}{CITE_DELIM}{locator}" + return f"{payload}{CITE_STOP}" + + +def _has_marker_codepoints(text: str) -> bool: + return any(c in text for c in (CITE_START, CITE_STOP, CITE_DELIM)) + + +def test_passthrough_when_no_marker_present(): + text = "Plain text with no citation markers." + assert _replace_openai_citation_markers(text, []) == text + + +def test_marker_rewritten_to_link_when_annotation_known(): + text = f"The capital is Paris {_marker('turn0view0')}." + citations = [ + { + "source_id": "turn0view0", + "url": "https://example.com/paris", + "title": "Paris", + }, + ] + out = _replace_openai_citation_markers(text, citations) + assert not _has_marker_codepoints(out) + assert "[[1]](https://example.com/paris)" in out + + +def test_unknown_source_marker_dropped_silently(): + text = f"Foo {_marker('turn9view9')} bar." + out = _replace_openai_citation_markers(text, []) + # Marker stripped, no garbled "E202" glyph leaks through, and the + # surrounding text stays intact. + assert not _has_marker_codepoints(out) + assert "E202" not in out + assert "turn9view9" not in out + assert "Foo" in out and "bar" in out + + +def test_multiple_concatenated_markers_resolved_in_order(): + """Real-world wire shape: a string of markers butted up against each other + after a sentence, as in the user-reported bug.""" + markers = "".join(_marker(f"turn{i}view{j}") for i, j in [(1, 0), (1, 1), (3, 0)]) + text = f"All animals ranked. {markers}" + citations = [ + {"source_id": "turn1view0", "url": "https://a.example/dog", "title": "Dog"}, + {"source_id": "turn1view1", "url": "https://a.example/cat", "title": "Cat"}, + {"source_id": "turn3view0", "url": "https://a.example/tiger", "title": "Tiger"}, + ] + out = _replace_openai_citation_markers(text, citations) + assert "[[1]](https://a.example/dog)" in out + assert "[[2]](https://a.example/cat)" in out + assert "[[3]](https://a.example/tiger)" in out + assert not _has_marker_codepoints(out) + + +def test_marker_with_locator_resolves(): + text = f"See {_marker('turn2file0', 'L8-L13')}." + citations = [ + {"source_id": "turn2file0", "url": "https://example.com/doc.txt"}, + ] + out = _replace_openai_citation_markers(text, citations) + assert "[[1]](https://example.com/doc.txt)" in out + assert "L8-L13" not in out # locator detail dropped; we just link. + assert not _has_marker_codepoints(out) + + +def test_mixed_known_and_unknown_markers(): + known = _marker("turn0view0") + unknown = _marker("turn0view99") + text = f"Known {known} and unknown {unknown}." + citations = [ + {"source_id": "turn0view0", "url": "https://example.com/known"}, + ] + out = _replace_openai_citation_markers(text, citations) + assert "[[1]](https://example.com/known)" in out + # Unknown markers leave no trace, but surrounding prose stays. + assert "Known" in out and "unknown" in out + assert not _has_marker_codepoints(out) + assert "E202" not in out + + +def test_empty_text_returns_verbatim(): + assert _replace_openai_citation_markers("", []) == "" + + +def test_idempotent_on_pre_stripped_text(): + """Pre-stripped text (no private-use codepoints) returns verbatim.""" + text = "citeturn1view0 plain" + assert _replace_openai_citation_markers(text, []) == text + + +@pytest.mark.parametrize( + "citation", + [ + {"url": "https://example.com/a"}, # no source_id at all + {"source_id": None, "url": "https://example.com/b"}, + {"source_id": "", "url": "https://example.com/c"}, + ], +) +def test_citation_without_source_id_does_not_crash(citation): + text = f"X {_marker('turnXviewY')} Y" + out = _replace_openai_citation_markers(text, [citation]) + # No mapping, marker stripped. Crash-free is the contract. + assert not _has_marker_codepoints(out) + assert "turnXviewY" not in out + + +def test_multiple_source_id_aliases_resolve_to_same_url(): + """Every alias for the same URL must resolve, not just the first. + Regression for the Codex P1 on the original PR.""" + a = _marker("turn0view0") + b = _marker("turn0view0_span_1") + c = _marker("turn0view0_span_2") + text = f"Triple {a}{b}{c} cite." + citations = [ + { + "source_ids": ["turn0view0", "turn0view0_span_1", "turn0view0_span_2"], + "url": "https://example.com/paris", + "title": "Paris", + }, + ] + out = _replace_openai_citation_markers(text, citations) + # All three aliases collapse onto citation [1] -- the URL is the + # same so it would be misleading to show three different numbers. + assert out.count("[[1]](https://example.com/paris)") == 3 + assert not _has_marker_codepoints(out) + + +def test_source_ids_list_and_legacy_source_id_both_resolve(): + """Mixed-shape citation: legacy ``source_id`` plus newer + ``source_ids`` aliases both resolve.""" + legacy = _marker("legacy_id") + alias = _marker("alias_id") + text = f"Both {legacy} and {alias} work." + citations = [ + { + "source_id": "legacy_id", + "source_ids": ["alias_id"], + "url": "https://example.com/doc", + }, + ] + out = _replace_openai_citation_markers(text, citations) + assert out.count("[[1]](https://example.com/doc)") == 2 + assert not _has_marker_codepoints(out) + + +# --------------------------------------------------------------------------- +# _rewrite_citation_markers_partial: deferred-annotation tests. OpenAI emits +# url_citation annotations on a subsequent SSE event; this helper reports +# `has_unresolved` so the stream loop defers emission. See PR #5713 audit. +# --------------------------------------------------------------------------- + + +def test_partial_known_marker_resolves_and_clears_unresolved(): + text = f"Foo {_marker('s1')} bar." + out, unresolved = _rewrite_citation_markers_partial( + text, + [{"source_id": "s1", "url": "https://example.com/a"}], + ) + assert "[[1]](https://example.com/a)" in out + assert unresolved is False + assert not _has_marker_codepoints(out) + + +def test_partial_unknown_marker_preserves_verbatim_and_flags(): + text = f"Foo {_marker('s1')} bar." + out, unresolved = _rewrite_citation_markers_partial(text, []) + assert unresolved is True + # Codepoints must remain so a follow-up pass can re-parse. + assert _has_marker_codepoints(out) + assert "Foo" in out and "bar." in out + + +def test_partial_resolves_after_late_annotation(): + """Two-pass: first call sees no citations, second resolves after annotation.""" + text = f"See {_marker('s1')} for details." + out1, unresolved1 = _rewrite_citation_markers_partial(text, []) + assert unresolved1 is True + citations = [{"source_id": "s1", "url": "https://example.com/x"}] + out2, unresolved2 = _rewrite_citation_markers_partial(out1, citations) + assert unresolved2 is False + assert "[[1]](https://example.com/x)" in out2 + assert not _has_marker_codepoints(out2) + + +def test_partial_multi_source_partial_resolution_keeps_marker_pending(): + """Any unresolved token in a multi-source marker leaves the whole marker + verbatim with ``unresolved`` True; defer until every id resolves or + end-of-stream forces a flush (dropping unresolved tokens then).""" + cite = f"{CITE_START}cite{CITE_DELIM}known{CITE_DELIM}locator{CITE_STOP}" + text = f"Pre {cite} post." + citations = [{"source_id": "known", "url": "https://example.com/y"}] + out, unresolved = _rewrite_citation_markers_partial(text, citations) + assert unresolved is True + assert cite in out + # End-of-stream force flush: drop the unresolved token, keep the + # resolved link. The streamer routes pending segments through + # `_replace_openai_citation_markers` at force=True for this. + forced = _replace_openai_citation_markers(out, citations) + assert "[[1]](https://example.com/y)" in forced + assert "locator" not in forced + assert not _has_marker_codepoints(forced) + + +def test_partial_idempotent_on_marker_free_text(): + text = "Plain text." + out, unresolved = _rewrite_citation_markers_partial(text, []) + assert out == text + assert unresolved is False + + +def test_partial_mixed_known_and_pending_markers_flags_unresolved(): + known = _marker("known") + pending = _marker("pending") + text = f"{known} {pending}" + citations = [{"source_id": "known", "url": "https://example.com/k"}] + out, unresolved = _rewrite_citation_markers_partial(text, citations) + assert unresolved is True # the pending marker drives the flag + assert "[[1]](https://example.com/k)" in out + # The pending marker stays verbatim for the next pass. + assert CITE_START in out and "pending" in out diff --git a/studio/backend/tests/test_openai_citation_markers_edge.py b/studio/backend/tests/test_openai_citation_markers_edge.py new file mode 100644 index 0000000000..ffe8c6b6eb --- /dev/null +++ b/studio/backend/tests/test_openai_citation_markers_edge.py @@ -0,0 +1,413 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0 + +"""Edge-case tests for the OpenAI Responses citation marker rewriter. + +Covers multi-source markers, source+locator, marker SPLIT across SSE deltas, +unterminated tails at end-of-stream, multiple markers per delta, late +annotation ordering, and idempotency. + +Reference: https://developers.openai.com/api/docs/guides/citation-formatting +""" + +import importlib + + +# Streaming integration is exercised by ``_simulate_delta_stream`` further +# down, mirroring the head/buffer/flush dance from ``_stream_openai_responses``. +_module = importlib.import_module("core.inference.external_provider") +_replace_openai_citation_markers = _module._replace_openai_citation_markers +_split_pending_citation_tail = _module._split_pending_citation_tail + + +CITE_START = "" +CITE_STOP = "" +CITE_DELIM = "" + + +def _marker(*source_ids: str, locator: str | None = None) -> str: + """Build a ``\\ue200cite\\ue202[\\ue202...][\\ue202]\\ue201`` + marker. Accepts one or many ``source_ids`` plus an optional ``locator``.""" + payload = f"{CITE_START}cite{CITE_DELIM}" + CITE_DELIM.join(source_ids) + if locator: + payload = f"{payload}{CITE_DELIM}{locator}" + return f"{payload}{CITE_STOP}" + + +def _no_private_use(text: str) -> bool: + return all(c not in text for c in (CITE_START, CITE_STOP, CITE_DELIM)) + + +# Harness mirroring the head/pending-tail/flush dance in +# `_stream_openai_responses`, so streaming tests skip the httpx mock. +def _simulate_delta_stream( + deltas: list[str], + citations: list[dict], + *, + flush: bool = True, +) -> str: + pending = "" + emitted: list[str] = [] + for delta in deltas: + combined = pending + delta + head, pending = _split_pending_citation_tail(combined) + if head: + head = _replace_openai_citation_markers(head, citations) + if head: + emitted.append(head) + if flush and pending: + # Mirror `_flush_pending_marker_tail`: drop the tail entirely if no + # closing stop byte arrived; the literal ``cite`` would leak otherwise. + if CITE_STOP not in pending: + rendered = "" + else: + rendered = _replace_openai_citation_markers(pending, citations) + for ch in (CITE_START, CITE_STOP, CITE_DELIM): + rendered = rendered.replace(ch, "") + import re as _re + + rendered = _re.sub(r"^cite\S*", "", rendered) + if rendered: + emitted.append(rendered) + return "".join(emitted) + + +# --------------------------------------------------------------------------- +# 1. Multi-source markers per the OpenAI docs. +# --------------------------------------------------------------------------- + + +def test_multi_source_marker_all_resolve(): + """\\ue200cite\\ue202id1\\ue202id2\\ue202id3\\ue201 expands to three links + when every id is known. Earlier regex captured only id1 and dropped id2/id3.""" + text = f"All three: {_marker('id1', 'id2', 'id3')}" + citations = [ + {"source_id": "id1", "url": "https://example.com/1"}, + {"source_id": "id2", "url": "https://example.com/2"}, + {"source_id": "id3", "url": "https://example.com/3"}, + ] + out = _replace_openai_citation_markers(text, citations) + assert "[[1]](https://example.com/1)" in out + assert "[[2]](https://example.com/2)" in out + assert "[[3]](https://example.com/3)" in out + assert _no_private_use(out) + + +def test_multi_source_marker_partial_resolution(): + """Known ids render, unknown ids drop silently, no glyph leaks.""" + text = f"Mixed: {_marker('known', 'unknown', 'also_known')}" + citations = [ + {"source_id": "known", "url": "https://k.example"}, + {"source_id": "also_known", "url": "https://ak.example"}, + ] + out = _replace_openai_citation_markers(text, citations) + assert "[[1]](https://k.example)" in out + assert "[[2]](https://ak.example)" in out + assert "unknown" not in out + assert _no_private_use(out) + + +# --------------------------------------------------------------------------- +# 2. Source + locator: locator is dropped, link still resolves. +# --------------------------------------------------------------------------- + + +def test_marker_with_numeric_locator(): + text = f"See {_marker('tu0', locator = '42')}." + citations = [{"source_id": "tu0", "url": "https://example.com/doc"}] + out = _replace_openai_citation_markers(text, citations) + assert "[[1]](https://example.com/doc)" in out + assert "42" not in out + assert _no_private_use(out) + + +def test_marker_with_range_locator(): + text = f"See {_marker('tu0', locator = 'L8-L13')}." + citations = [{"source_id": "tu0", "url": "https://example.com/code"}] + out = _replace_openai_citation_markers(text, citations) + assert "[[1]](https://example.com/code)" in out + assert "L8-L13" not in out + assert _no_private_use(out) + + +# --------------------------------------------------------------------------- +# 3. Marker SPLIT across two SSE deltas -- the codex-flagged P1. +# --------------------------------------------------------------------------- + + +def test_marker_split_in_source_id(): + """Delta-1 ends mid-source-id (``\\ue200cite\\ue202tu``), delta-2 starts + with the rest (``rn0view0\\ue201``). The buffer stitches the halves + back together so they resolve to one link instead of leaking.""" + full = f"See {_marker('turn0view0')} now." + # Cut right after the second delim + "tu" inside the source id. + cut = full.index("tu", full.index(CITE_START)) + len("tu") + d1, d2 = full[:cut], full[cut:] + # Sanity check: delta-1 actually contains a partial marker. + assert CITE_START in d1 and CITE_STOP not in d1 + assert CITE_STOP in d2 + citations = [{"source_id": "turn0view0", "url": "https://x"}] + out = _simulate_delta_stream([d1, d2], citations) + assert out == "See [[1]](https://x) now." + assert _no_private_use(out) + + +def test_marker_split_at_start_byte(): + """Split exactly after the opening ``\\ue200`` byte; the buffer must + hold the lone open byte until the rest arrives.""" + full = f"Text {_marker('sid')} done" + cut = full.index(CITE_START) + 1 # right AFTER the open byte + d1, d2 = full[:cut], full[cut:] + citations = [{"source_id": "sid", "url": "https://y"}] + out = _simulate_delta_stream([d1, d2], citations) + assert out == "Text [[1]](https://y) done" + assert _no_private_use(out) + + +def test_marker_split_across_three_deltas(): + """Worst case: marker chopped into three pieces across three deltas.""" + full = f"A {_marker('threesplit')} B" + # cut at two points inside the marker + open_pos = full.index(CITE_START) + stop_pos = full.index(CITE_STOP) + cut1 = open_pos + 4 + cut2 = stop_pos - 2 + parts = [full[:cut1], full[cut1:cut2], full[cut2:]] + citations = [{"source_id": "threesplit", "url": "https://z"}] + out = _simulate_delta_stream(parts, citations) + assert out == "A [[1]](https://z) B" + assert _no_private_use(out) + + +def test_marker_split_with_trailing_text_after_close(): + """Delta-2 closes the marker AND carries trailing prose; both emit cleanly.""" + full = f"X {_marker('sid')} after" + cut = full.index("cite") + len("ci") + d1, d2 = full[:cut], full[cut:] + citations = [{"source_id": "sid", "url": "https://a"}] + out = _simulate_delta_stream([d1, d2], citations) + assert out == "X [[1]](https://a) after" + assert _no_private_use(out) + + +def test_split_marker_unknown_source_is_dropped_cleanly(): + """Split marker for an unknown source drops silently on flush.""" + full = f"Pre {_marker('never_seen')} post" + cut = full.index(CITE_START) + 3 + d1, d2 = full[:cut], full[cut:] + out = _simulate_delta_stream([d1, d2], []) + assert out == "Pre post" + assert _no_private_use(out) + + +# --------------------------------------------------------------------------- +# 4. Unterminated marker at end-of-stream -- truncation safety. +# --------------------------------------------------------------------------- + + +def test_unterminated_marker_at_stream_end_dropped_on_flush(): + """Stream ends mid-marker (e.g. response.incomplete); the tail is + flushed with private-use bytes stripped, no `E202` text leaks.""" + deltas = ["Some text ", f"{CITE_START}citetu", "rn0view0"] # no STOP ever + out = _simulate_delta_stream(deltas, [], flush = True) + assert _no_private_use(out) + assert "E200" not in out and "E202" not in out + # Surrounding prose stays; we don't assert exact marker remainder. + assert "Some text " in out + + +def test_flush_resolves_marker_when_late_annotation_arrives(): + """Marker in a delta, matching annotation arrives later (on + response.output_text.annotation.added after the final delta). The + rewriter reads ``all_url_citations`` LIVE at flush, so the buffered + marker still resolves.""" + deltas = ["Look ", f"{CITE_START}cite{CITE_DELIM}late_sid"] + pending = "" + citations: list[dict] = [] + emitted: list[str] = [] + for d in deltas: + combined = pending + d + head, pending = _split_pending_citation_tail(combined) + if head: + emitted.append(_replace_openai_citation_markers(head, citations)) + # Annotation arrives AFTER all deltas but BEFORE flush. + citations.append({"source_id": "late_sid", "url": "https://late.example"}) + # Append the STOP byte that closed the marker in a later delta. + pending = pending + CITE_STOP + flushed = _replace_openai_citation_markers(pending, citations) + for ch in (CITE_START, CITE_STOP, CITE_DELIM): + flushed = flushed.replace(ch, "") + emitted.append(flushed) + out = "".join(emitted) + assert "[[1]](https://late.example)" in out + assert _no_private_use(out) + + +# --------------------------------------------------------------------------- +# 5. Multiple unrelated markers in a single delta. +# --------------------------------------------------------------------------- + + +def test_three_markers_in_one_delta_resolve_independently(): + text = f"alpha {_marker('a')} beta {_marker('b')} gamma {_marker('c')} end" + citations = [ + {"source_id": "a", "url": "https://example.com/a"}, + {"source_id": "b", "url": "https://example.com/b"}, + {"source_id": "c", "url": "https://example.com/c"}, + ] + out = _replace_openai_citation_markers(text, citations) + assert out == ( + "alpha [[1]](https://example.com/a) beta " + "[[2]](https://example.com/b) gamma " + "[[3]](https://example.com/c) end" + ) + + +# --------------------------------------------------------------------------- +# 6. Idempotency. +# --------------------------------------------------------------------------- + + +def test_rewriter_idempotent_on_already_rewritten_text(): + """Running the rewriter twice does not double-link or corrupt brackets.""" + text = f"alpha {_marker('a')} omega" + citations = [{"source_id": "a", "url": "https://example.com/a"}] + once = _replace_openai_citation_markers(text, citations) + twice = _replace_openai_citation_markers(once, citations) + assert once == twice + assert _no_private_use(once) + + +def test_rewriter_idempotent_on_marker_free_text(): + """No-op when there is nothing to rewrite.""" + text = "Plain prose with no citations and no private-use bytes." + out = _replace_openai_citation_markers(text, []) + assert out is text or out == text + + +# --------------------------------------------------------------------------- +# 7. Edge / robustness. +# --------------------------------------------------------------------------- + + +def test_only_marker_no_surrounding_text(): + """A delta that is JUST a marker (no prose) still renders correctly; + used to leak without the empty-string short-circuit in the split helper.""" + text = _marker("solo") + citations = [{"source_id": "solo", "url": "https://solo.example"}] + out = _replace_openai_citation_markers(text, citations) + assert out == "[[1]](https://solo.example)" + + +def test_back_to_back_markers_with_no_separator(): + """Adjacent markers resolve to concatenated links, no joining whitespace.""" + text = f"{_marker('x')}{_marker('y')}" + citations = [ + {"source_id": "x", "url": "https://x.example"}, + {"source_id": "y", "url": "https://y.example"}, + ] + out = _replace_openai_citation_markers(text, citations) + assert out == "[[1]](https://x.example)[[2]](https://y.example)" + + +def test_split_helper_buffers_only_after_last_open_byte(): + """A complete marker followed by an unterminated one: head includes + the complete marker, buffer holds only the trailing partial.""" + complete = _marker("done") + partial = f"{CITE_START}cite{CITE_DELIM}half" # no STOP + text = f"pre {complete} mid {partial}" + head, tail = _split_pending_citation_tail(text) + assert head == f"pre {complete} mid " + assert tail == partial + # And the head, once rewritten, drops every private-use byte. + rewritten = _replace_openai_citation_markers( + head, [{"source_id": "done", "url": "https://d"}] + ) + assert rewritten == "pre [[1]](https://d) mid " + + +def test_split_helper_empty_input(): + head, tail = _split_pending_citation_tail("") + assert head == "" and tail == "" + + +def test_split_helper_no_open_byte(): + head, tail = _split_pending_citation_tail("nothing to see here") + assert head == "nothing to see here" and tail == "" + + +def test_split_helper_complete_marker_only(): + """A delta ending with a closed marker leaves the buffer empty.""" + text = f"alpha {_marker('a')}" + head, tail = _split_pending_citation_tail(text) + assert head == text and tail == "" + + +# --------------------------------------------------------------------------- +# 8. Sources-panel: marker drop must not affect citation aggregation. +# Indices come from the url_citations list, not the marker stream. +# --------------------------------------------------------------------------- + + +def test_unknown_marker_does_not_perturb_citation_indexing(): + """Unknown source_id markers drop without consuming an index slot.""" + text = f"A {_marker('unknown')} B {_marker('real_a')} C {_marker('real_b')}" + citations = [ + {"source_id": "real_a", "url": "https://example.com/a"}, + {"source_id": "real_b", "url": "https://example.com/b"}, + ] + out = _replace_openai_citation_markers(text, citations) + # real_a is index 1; unknown does not take a slot. + assert "[[1]](https://example.com/a)" in out + assert "[[2]](https://example.com/b)" in out + assert _no_private_use(out) + + +# --------------------------------------------------------------------------- +# Regression: unterminated marker tail must NOT leak the residual +# ``cite``-prefixed source id as plain text. PR #5713 audit P1. +# --------------------------------------------------------------------------- + + +def test_unterminated_marker_does_not_leak_cite_residue(): + """Stream ends mid-marker: drop the whole tail rather than strip + codepoints and leave ``cite`` behind.""" + half = f"Hi there {CITE_START}cite{CITE_DELIM}turn0view0" + out = _simulate_delta_stream([half], [], flush = True) + # Prose before the marker stays; no private-use bytes or cite residue. + assert "Hi there" in out + assert _no_private_use(out) + assert "citeturn0view0" not in out + assert "cite" not in out.split("Hi there", 1)[1] + + +def test_unterminated_marker_only_no_prefix_drops_entirely(): + """A delta that is purely an unterminated marker flushes to "".""" + half = f"{CITE_START}cite{CITE_DELIM}turn0view0" + out = _simulate_delta_stream([half], [], flush = True) + assert out == "" + + +def test_unterminated_marker_with_prefix_emits_only_prefix(): + """Prose then unterminated marker: prose emits, marker remnant drops.""" + half = f"prefix prose {CITE_START}cite{CITE_DELIM}abc" + out = _simulate_delta_stream([half], [], flush = True) + assert out == "prefix prose " + + +def test_closing_byte_arrives_after_pending_buffered_split(): + """Closing byte arrives in a later delta after opener + source id were + buffered; link resolves with no residue.""" + cuts = [ + f"a {CITE_START}cite{CITE_DELIM}", + f"sid{CITE_STOP} b", + ] + out = _simulate_delta_stream( + cuts, + [{"source_id": "sid", "url": "https://example.com/x"}], + flush = True, + ) + assert "[[1]](https://example.com/x)" in out + assert "a " in out and "b" in out + assert _no_private_use(out) + assert "citesid" not in out diff --git a/studio/backend/tests/test_pricing.py b/studio/backend/tests/test_pricing.py index cc8c16993c..313c15a441 100644 --- a/studio/backend/tests/test_pricing.py +++ b/studio/backend/tests/test_pricing.py @@ -1,12 +1,8 @@ # SPDX-License-Identifier: AGPL-3.0-only # Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0 -"""Unit tests for the per-session cost calculator. - -Pricing inputs are baked into ``core/inference/pricing.py``; this -test verifies the math (with multipliers from the prompt-caching -docs) and that unknown models / empty usage degrade gracefully. -""" +"""Unit tests for the per-session cost calculator. Verifies math +against ``core/inference/pricing.py`` and graceful degradation.""" import math @@ -14,6 +10,7 @@ from core.inference.pricing import ( ANTHROPIC_CACHE_5M_WRITE_MULT, ANTHROPIC_CACHE_1H_WRITE_MULT, ANTHROPIC_CACHE_READ_MULT, + ANTHROPIC_FAST_MODE_MULT, ANTHROPIC_PRICING, OPENAI_CACHE_READ_MULT, OPENAI_CONTAINER_USD_PER_HOUR, @@ -57,6 +54,64 @@ def test_anthropic_opus_4_7_input_and_output_math(): assert _isclose(out["total_usd"], 30.0) +# ── Anthropic fast-mode 6x multiplier (Opus 4.6 / 4.7 only) ───────── + + +def test_anthropic_fast_mode_charges_6x_standard_opus(): + """6x on input + output when ``usage.speed == "fast"``. + https://platform.claude.com/docs/en/build-with-claude/fast-mode""" + out = calculate_cost( + "anthropic", + "claude-opus-4-7", + { + "input_tokens": 1_000_000, + "output_tokens": 1_000_000, + "speed": "fast", + }, + ) + assert _isclose(out["input_usd"], 5.0 * ANTHROPIC_FAST_MODE_MULT) + assert _isclose(out["output_usd"], 25.0 * ANTHROPIC_FAST_MODE_MULT) + assert _isclose(out["total_usd"], 30.0 * ANTHROPIC_FAST_MODE_MULT) + assert "(fast)" in out["model_priced"], out["model_priced"] + + +def test_anthropic_fast_mode_does_not_affect_standard_speed(): + """``speed: "standard"`` (or missing) keeps the base rates.""" + out_standard = calculate_cost( + "anthropic", + "claude-opus-4-7", + { + "input_tokens": 1_000_000, + "output_tokens": 1_000_000, + "speed": "standard", + }, + ) + out_missing = calculate_cost( + "anthropic", + "claude-opus-4-7", + {"input_tokens": 1_000_000, "output_tokens": 1_000_000}, + ) + assert _isclose(out_standard["total_usd"], out_missing["total_usd"]) + assert _isclose(out_standard["total_usd"], 30.0) + + +def test_anthropic_fast_mode_stacks_with_cache_read_multiplier(): + """Cache multipliers apply on top of fast-mode (per docs).""" + base = ANTHROPIC_PRICING["claude-opus-4-7"]["input_per_mtok"] + out = calculate_cost( + "anthropic", + "claude-opus-4-7", + { + "input_tokens": 0, + "output_tokens": 0, + "cache_read_input_tokens": 1_000_000, + "speed": "fast", + }, + ) + expected = base * ANTHROPIC_FAST_MODE_MULT * ANTHROPIC_CACHE_READ_MULT + assert _isclose(out["cache_read_usd"], expected) + + # ── Anthropic cache write 5m + read multipliers ────────────────────── @@ -102,8 +157,7 @@ def test_anthropic_cache_1h_write_uses_2x_multiplier(): def test_anthropic_cache_5m_default_when_no_breakdown(): - # When the docs/response doesn't surface the 5m/1h split, treat - # the full cache_creation bucket as 5m (the upstream default pool). + # No 5m/1h split surfaced -> assume the default 5m pool. base = ANTHROPIC_PRICING["claude-opus-4-7"]["input_per_mtok"] out = calculate_cost( "anthropic", @@ -148,8 +202,7 @@ def test_anthropic_code_exec_charged_per_hour(): def test_anthropic_dated_id_falls_back_to_canonical_prefix(): - # Hypothetical dated snapshot of claude-opus-4-7 should still - # inherit the canonical-id pricing via the prefix-match fallback. + # Dated snapshot inherits canonical pricing via prefix-match. out = calculate_cost( "anthropic", "claude-opus-4-7-20260712", @@ -163,8 +216,7 @@ def test_anthropic_dated_id_falls_back_to_canonical_prefix(): def test_openai_gpt55_input_output_math(): - # Sub-272k input keeps us in the short-context tier ($5/$30). - # The dedicated long-context tests below exercise the crossover. + # Sub-272k stays in short-context tier ($5/$30). out = calculate_cost( "openai", "gpt-5.5", @@ -176,11 +228,7 @@ def test_openai_gpt55_input_output_math(): def test_openai_cache_read_subtracted_from_input_at_discount(): - # OpenAI folds cached tokens into input_tokens, unlike Anthropic. - # The calculator must subtract cached_tokens from the "full price" - # bucket and re-bill them at 0.1x. Use a sub-272k total so the - # short-context tier applies (long-context crossover is exercised - # in its own test below). + # OpenAI folds cached into input_tokens; subtract and re-bill at 0.1x. base = OPENAI_PRICING["gpt-5.5"]["input_per_mtok"] out = calculate_cost( "openai", @@ -199,9 +247,7 @@ def test_openai_cache_read_subtracted_from_input_at_discount(): def test_openai_billable_input_tokens_does_not_double_count_cache_read(): - # OpenAI's input_tokens already includes cached_tokens, so the - # billable counter must NOT add cache_read on top -- otherwise the - # tooltip says 180k input when the bill is for 100k. + # input_tokens already includes cached; don't double-count. out = calculate_cost( "openai", "gpt-5.5", @@ -215,9 +261,7 @@ def test_openai_billable_input_tokens_does_not_double_count_cache_read(): def test_openai_dated_snapshot_inherits_canonical_pricing(): - # Sub-272k stays in the short-context tier; the prefix-match - # fallback is what proves the dated snapshot inherits gpt-5.5 - # pricing. + # Dated snapshot inherits gpt-5.5 pricing via prefix-match. out = calculate_cost( "openai", "gpt-5.5-2026-04-23", @@ -228,10 +272,7 @@ def test_openai_dated_snapshot_inherits_canonical_pricing(): def test_openai_gpt54_family_uses_verified_prices(): - # Spot-check the lower-tier rows that previously underbilled. - # gpt-5.4 has a long-context tier so the input has to stay - # below 272k; the mini/nano/codex rows have no crossover so - # 1M tokens is fine. + # Spot-check lower-tier rows that previously underbilled. cases = { # (input_tokens, expected_input_usd, expected_output_usd) "gpt-5.4": (200_000, 200_000 / 1_000_000.0 * 2.5, 200_000 / 1_000_000.0 * 15.0), @@ -251,8 +292,7 @@ def test_openai_gpt54_family_uses_verified_prices(): def test_openai_unlisted_model_priced_false_not_zero_default(): - # o-series / gpt-4.5 are no longer on the pricing page, so we - # intentionally drop them rather than silently underbill at $0. + # o-series / gpt-4.5 are off the pricing page; drop rather than $0. for model in ("o3", "o4-mini", "gpt-4.5", "gpt-4.5-preview"): out = calculate_cost( "openai", @@ -270,9 +310,7 @@ def test_openai_unlisted_model_priced_false_not_zero_default(): def test_anthropic_canonical_4_5_ids_are_priced(): - # Codex P1: claude-opus-4-5 (no date) is the canonical id used - # in backend defaults but was missing from the table, so the - # calculator returned priced=False + zero cost. Pin the aliases. + # Pin the bare-id aliases (backend defaults reference these). cases = { "claude-opus-4-5": (5.0, 25.0), "claude-sonnet-4-5": (3.0, 15.0), @@ -307,8 +345,7 @@ def test_openai_gpt55_short_context_under_272k_uses_base_rates(): def test_openai_gpt55_long_context_crossover_uses_higher_rates(): - # 300k billable input > 272k threshold -> long-context tier - # applies to the WHOLE turn, not a per-token blend. + # >272k billable -> long-context tier on the whole turn. out = calculate_cost( "openai", "gpt-5.5", @@ -330,8 +367,7 @@ def test_openai_gpt54_long_context_crossover(): def test_openai_gpt54_mini_has_no_long_context_tier(): - # Mini/nano/codex don't publish a long-context price; the base - # rate must keep applying even at very large prompts. + # Mini/nano/codex have no long-context tier; base rate always applies. out = calculate_cost( "openai", "gpt-5.4-mini", @@ -374,8 +410,7 @@ def test_openai_container_hours_charged(): def test_openai_tool_surcharges_added_to_total(): - # End-to-end: input + output + web_search + container in one - # turn. Total must sum all four buckets. + # End-to-end: total must sum input + output + web_search + container. out = calculate_cost( "openai", "gpt-5.5", @@ -412,12 +447,12 @@ def test_snapshot_contains_provider_buckets_and_multipliers(): assert a["cache_5m_write_mult"] == ANTHROPIC_CACHE_5M_WRITE_MULT assert a["cache_1h_write_mult"] == ANTHROPIC_CACHE_1H_WRITE_MULT assert a["cache_read_mult"] == ANTHROPIC_CACHE_READ_MULT + assert a["fast_mode_mult"] == ANTHROPIC_FAST_MODE_MULT assert "web_search_usd_per_1k" in a assert "code_execution_usd_per_hour" in a assert "models" in o and "gpt-5.5" in o["models"] assert o["cache_read_mult"] == OPENAI_CACHE_READ_MULT - # OpenAI tool surcharge constants are also exposed so the frontend - # tooltip can render the per-call rate. + # OpenAI tool surcharge constants are exposed for the frontend. assert o["web_search_usd_per_1k"] == OPENAI_WEB_SEARCH_USD_PER_1K assert o["container_usd_per_hour"] == OPENAI_CONTAINER_USD_PER_HOUR # Long-context tier metadata travels with the model row. @@ -425,3 +460,169 @@ def test_snapshot_contains_provider_buckets_and_multipliers(): assert gpt55["long_context_threshold"] == 272_000 assert gpt55["long_context_input_per_mtok"] == 10.0 assert gpt55["long_context_output_per_mtok"] == 45.0 + + +# ── longest-prefix match: dated mini variant must not collide with the +# shorter family prefix. ── + + +def test_longest_prefix_match_wins_for_dated_mini_snapshot(): + """`gpt-5.4-mini-2026-...` must inherit the mini rate, not the + shorter `gpt-5.4` rate (longest prefix wins).""" + out = calculate_cost( + "openai", + "gpt-5.4-mini-2026-04-23", + {"input_tokens": 1_000_000, "output_tokens": 0}, + ) + assert out["priced"] is True + # mini = 0.75/MTok, shorter gpt-5.4 = 2.5/MTok (>3x overcharge). + assert _isclose(out["input_usd"], 0.75), out + + +def test_longest_prefix_match_wins_for_dated_pro_snapshot(): + out = calculate_cost( + "openai", + "gpt-5.5-pro-2026-04-23", + {"input_tokens": 1_000_000, "output_tokens": 0}, + ) + assert out["priced"] is True + # gpt-5.5-pro = 30/MTok vs gpt-5.5 = 5/MTok; longest wins. + assert _isclose(out["input_usd"], 30.0), out + + +# ── accept both chat-style and Responses envelope shapes. ── + + +def test_openai_chat_style_usage_keys_priced_correctly(): + """Chat-style envelope (`prompt_tokens` / `completion_tokens`) must + produce a non-zero cost (previously silently zeroed).""" + out = calculate_cost( + "openai", + "gpt-5.4-mini", + {"prompt_tokens": 1_000_000, "completion_tokens": 1_000_000}, + ) + # gpt-5.4-mini: 0.75 input + 4.5 output per MTok. + assert _isclose(out["input_usd"], 0.75), out + assert _isclose(out["output_usd"], 4.5), out + + +def test_input_tokens_preferred_when_both_keys_present(): + """Raw key wins when both envelope shapes are present.""" + out = calculate_cost( + "openai", + "gpt-5.4-mini", + { + "input_tokens": 2_000_000, + "prompt_tokens": 5_000_000, + "output_tokens": 0, + }, + ) + # input_tokens=2M wins -> 2 * 0.75 = 1.50. + assert _isclose(out["input_usd"], 1.50), out + + +def test_anthropic_chat_style_prompt_tokens_dedupes_cache_buckets(): + """Anthropic chat-style prompt_tokens already folds cache buckets; + don't double-count billable input.""" + # 1M uncached + 200K cache_creation + 500K cache_read -> 1.7M folded. + raw = calculate_cost( + "anthropic", + "claude-opus-4-7", + { + "input_tokens": 1_000_000, + "cache_creation_input_tokens": 200_000, + "cache_read_input_tokens": 500_000, + "output_tokens": 0, + }, + ) + chat = calculate_cost( + "anthropic", + "claude-opus-4-7", + { + "prompt_tokens": 1_700_000, + "cache_creation_input_tokens": 200_000, + "cache_read_input_tokens": 500_000, + "completion_tokens": 0, + }, + ) + # Both envelopes must price the same. + assert _isclose(chat["input_usd"], raw["input_usd"]), (chat, raw) + assert _isclose(chat["cache_write_usd"], raw["cache_write_usd"]), (chat, raw) + assert _isclose(chat["cache_read_usd"], raw["cache_read_usd"]), (chat, raw) + assert _isclose(chat["total_usd"], raw["total_usd"]), (chat, raw) + assert chat["billable_input_tokens"] == raw["billable_input_tokens"], (chat, raw) + + +def test_openai_chat_style_prompt_tokens_keeps_cache_read_semantics(): + """OpenAI prompt_tokens includes cache_read like raw input_tokens.""" + raw = calculate_cost( + "openai", + "gpt-5.5", + { + "input_tokens": 1_000_000, + "input_tokens_details": {"cached_tokens": 200_000}, + "output_tokens": 100_000, + }, + ) + chat = calculate_cost( + "openai", + "gpt-5.5", + { + "prompt_tokens": 1_000_000, + "cache_read_input_tokens": 200_000, + "completion_tokens": 100_000, + }, + ) + assert _isclose(chat["total_usd"], raw["total_usd"]), (chat, raw) + + +def test_openai_chat_style_envelope_reads_cache_from_prompt_tokens_details(): + """Chat-style envelope ships cached under prompt_tokens_details; + calculator must honour both this and input_tokens_details.""" + base = OPENAI_PRICING["gpt-5.5"]["input_per_mtok"] + raw = calculate_cost( + "openai", + "gpt-5.5", + { + "input_tokens": 100_000, + "input_tokens_details": {"cached_tokens": 80_000}, + "output_tokens": 0, + }, + ) + chat_style = calculate_cost( + "openai", + "gpt-5.5", + { + "prompt_tokens": 100_000, + "prompt_tokens_details": {"cached_tokens": 80_000}, + "completion_tokens": 0, + }, + ) + # Both envelopes must price identically. + assert _isclose(chat_style["input_usd"], raw["input_usd"]), (chat_style, raw) + assert _isclose(chat_style["cache_read_usd"], raw["cache_read_usd"]), ( + chat_style, + raw, + ) + # 80k at 0.1x base, 20k at full. + assert _isclose( + chat_style["cache_read_usd"], + 80_000 / 1_000_000.0 * base * OPENAI_CACHE_READ_MULT, + ) + + +def test_explicit_zero_output_tokens_wins_over_stale_completion_tokens(): + """Explicit ``output_tokens: 0`` beats a stale ``completion_tokens``; + the previous `or` fallback treated 0 as missing.""" + out = calculate_cost( + "openai", + "gpt-4o-mini", + { + "input_tokens": 100, + "output_tokens": 0, + # Stale chat-style mirror; must not bill against it. + "completion_tokens": 50, + }, + ) + assert out["billable_output_tokens"] == 0, out + assert out["output_usd"] == 0.0, out diff --git a/studio/backend/tests/test_pricing_edge.py b/studio/backend/tests/test_pricing_edge.py new file mode 100644 index 0000000000..ca4be258e0 --- /dev/null +++ b/studio/backend/tests/test_pricing_edge.py @@ -0,0 +1,475 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0 + +"""Adversarial edge cases for ``calculate_cost`` / ``_lookup``: prefix +boundary, negative tokens, chat vs raw parity, long-context crossover +on billable count, and malformed sub-objects.""" + +import math + +from core.inference.pricing import ( + ANTHROPIC_CACHE_5M_WRITE_MULT, + ANTHROPIC_CACHE_READ_MULT, + ANTHROPIC_PRICING, + OPENAI_CACHE_READ_MULT, + OPENAI_PRICING, + _lookup, + calculate_cost, +) + + +def _isclose(a, b, tol = 1e-6): + return math.isclose(a, b, rel_tol = tol, abs_tol = tol) + + +# ── prefix-match boundary checks ──────────────────────────────────── + + +def test_prefix_match_requires_dash_boundary_opus_variant(): + # `claude-opus-4-15` must not inherit `claude-opus-4-1` pricing; + # next char must be `-` or end-of-string. + assert _lookup("anthropic", "claude-opus-4-15") is None + out = calculate_cost( + "anthropic", + "claude-opus-4-15", + {"input_tokens": 1_000_000, "output_tokens": 0}, + ) + assert out["priced"] is False + assert out["total_usd"] == 0.0 + + +def test_prefix_match_requires_dash_boundary_gpt_variant(): + # Same dash-boundary invariant for OpenAI ids. + assert _lookup("openai", "gpt-5.55") is None + assert _lookup("openai", "gpt-5.55-2026-04-23") is None + out = calculate_cost( + "openai", + "gpt-5.55-2026-04-23", + {"input_tokens": 1_000_000, "output_tokens": 0}, + ) + assert out["priced"] is False + + +def test_prefix_match_requires_dash_boundary_pro_lookalike(): + # `gpt-5.5-prod` must fall through `gpt-5.5-pro` (6x overcharge) + # and land on the canonical `gpt-5.5` row. + prices = _lookup("openai", "gpt-5.5-prod") + assert prices is not None + assert ( + prices["input_per_mtok"] == OPENAI_PRICING["gpt-5.5"]["input_per_mtok"] + ), "expected fallback to gpt-5.5 base ($5), not gpt-5.5-pro ($30)" + out = calculate_cost( + "openai", + "gpt-5.5-prod", + {"input_tokens": 100_000, "output_tokens": 0}, + ) + assert out["priced"] is True + assert _isclose(out["input_usd"], 100_000 / 1_000_000.0 * 5.0) + + +def test_prefix_match_still_resolves_legit_dated_snapshots(): + # Boundary fix must not regress legit dated snapshots. + out = calculate_cost( + "openai", + "gpt-5.4-mini-2026-04-23", + {"input_tokens": 1_000_000, "output_tokens": 0}, + ) + assert out["priced"] is True + assert _isclose(out["input_usd"], 0.75) + + # And Anthropic dated snapshot still resolves to canonical row. + out = calculate_cost( + "anthropic", + "claude-opus-4-7-20260414", + {"input_tokens": 1_000_000, "output_tokens": 0}, + ) + assert out["priced"] is True + assert _isclose(out["input_usd"], 5.0) + + +# ── precedence: input_tokens wins over prompt_tokens (and 0 is real) ── + + +def test_explicit_zero_input_tokens_wins_over_stale_prompt_tokens(): + # Input-side mirror of the output zero precedence test. + out = calculate_cost( + "openai", + "gpt-5.5", + { + "input_tokens": 0, + "prompt_tokens": 1_000_000, # stale chat-style mirror + "output_tokens": 100, + }, + ) + assert out["billable_input_tokens"] == 0 + assert out["input_usd"] == 0.0 + + +def test_none_input_tokens_falls_through_to_prompt_tokens(): + # `None` is "key present but unset"; chat-style mirror wins. + out = calculate_cost( + "openai", + "gpt-5.5", + { + "input_tokens": None, + "prompt_tokens": 200_000, + "output_tokens": None, + "completion_tokens": 5_000, + }, + ) + assert out["billable_input_tokens"] == 200_000 + assert out["billable_output_tokens"] == 5_000 + assert _isclose(out["input_usd"], 200_000 / 1_000_000.0 * 5.0) + assert _isclose(out["output_usd"], 5_000 / 1_000_000.0 * 30.0) + + +# ── negative / corrupted upstream values clamp to zero ────────────── + + +def test_negative_tokens_clamp_to_zero_no_negative_bill(): + out = calculate_cost( + "openai", + "gpt-5.5", + {"input_tokens": -100, "output_tokens": -50}, + ) + assert out["billable_input_tokens"] == 0 + assert out["billable_output_tokens"] == 0 + assert out["input_usd"] == 0.0 + assert out["output_usd"] == 0.0 + assert out["total_usd"] == 0.0 + + +def test_negative_cache_buckets_clamp_to_zero(): + # Negative cache_read on Anthropic would otherwise refund the bill. + out = calculate_cost( + "anthropic", + "claude-opus-4-7", + { + "input_tokens": 1_000, + "output_tokens": 0, + "cache_creation_input_tokens": -500, + "cache_read_input_tokens": -1_000, + }, + ) + assert out["cache_write_usd"] == 0.0 + assert out["cache_read_usd"] == 0.0 + assert out["billable_input_tokens"] == 1_000 + assert out["total_usd"] >= 0.0 + + +def test_negative_prompt_tokens_chat_style_clamp(): + out = calculate_cost( + "openai", + "gpt-5.4-mini", + {"prompt_tokens": -100, "completion_tokens": -50}, + ) + assert out["billable_input_tokens"] == 0 + assert out["billable_output_tokens"] == 0 + assert out["total_usd"] == 0.0 + + +# ── cache_read > prompt_tokens corruption: no negative billable ───── + + +def test_anthropic_chat_cache_read_exceeds_prompt_no_negative_billable(): + # cache_read > prompt_tokens clamps uncached_input at 0; billable + # still reflects cache buckets (we charge for what we got). + out = calculate_cost( + "anthropic", + "claude-opus-4-7", + { + "prompt_tokens": 100, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 500, + "completion_tokens": 0, + }, + ) + assert out["input_usd"] == 0.0 # uncached clamped to 0 + assert out["billable_input_tokens"] == 500 # 0 uncached + 500 cache_read + # cache_read still priced at the discount rate. + base = ANTHROPIC_PRICING["claude-opus-4-7"]["input_per_mtok"] + assert _isclose( + out["cache_read_usd"], 500 / 1_000_000.0 * base * ANTHROPIC_CACHE_READ_MULT + ) + + +def test_openai_raw_cached_tokens_exceeds_input_clamp_non_cached(): + # OpenAI variant: cached > input must not produce negative input_usd. + base = OPENAI_PRICING["gpt-5.5"]["input_per_mtok"] + out = calculate_cost( + "openai", + "gpt-5.5", + { + "input_tokens": 100, + "output_tokens": 0, + "input_tokens_details": {"cached_tokens": 500}, + }, + ) + assert out["input_usd"] == 0.0 + # Cache read still priced (the 0.1x bucket). + assert _isclose( + out["cache_read_usd"], 500 / 1_000_000.0 * base * OPENAI_CACHE_READ_MULT + ) + + +# ── long-context tier crosses on billable, including cache_creation ── + + +def test_openai_long_context_triggers_on_cache_creation_inflated_billable(): + # cache_creation pushes billable past 272k -> long-context tier + # must fire to avoid undercounting. + out = calculate_cost( + "openai", + "gpt-5.5", + { + "input_tokens": 250_000, + "cache_creation_input_tokens": 50_000, + "output_tokens": 1_000, + }, + ) + assert out["billable_input_tokens"] == 300_000 + assert "long-context" in out["model_priced"] + assert _isclose(out["input_usd"], 250_000 / 1_000_000.0 * 10.0) + assert _isclose(out["output_usd"], 1_000 / 1_000_000.0 * 45.0) + + +def test_openai_long_context_threshold_boundary_inclusive(): + # Threshold is inclusive (>=). + out = calculate_cost( + "openai", + "gpt-5.5", + {"input_tokens": 272_000, "output_tokens": 1_000}, + ) + assert "long-context" in out["model_priced"] + out_lo = calculate_cost( + "openai", + "gpt-5.5", + {"input_tokens": 271_999, "output_tokens": 1_000}, + ) + assert "long-context" not in out_lo["model_priced"] + + +# ── chat-style vs raw envelope parity at OpenAI long-context tier ── + + +def test_openai_chat_envelope_long_context_parity_with_raw(): + raw = calculate_cost( + "openai", + "gpt-5.5", + {"input_tokens": 300_000, "output_tokens": 10_000}, + ) + chat = calculate_cost( + "openai", + "gpt-5.5", + {"prompt_tokens": 300_000, "completion_tokens": 10_000}, + ) + assert _isclose(chat["total_usd"], raw["total_usd"]) + assert "long-context" in chat["model_priced"] + assert "long-context" in raw["model_priced"] + + +# ── malformed sub-objects: no crash, no false bill ────────────────── + + +def test_cache_creation_as_int_does_not_crash(): + # Proxies sometimes fold cache_creation to an int; tolerate it + # and fall back to the 5m default. + base = ANTHROPIC_PRICING["claude-opus-4-7"]["input_per_mtok"] + out = calculate_cost( + "anthropic", + "claude-opus-4-7", + { + "input_tokens": 0, + "output_tokens": 0, + "cache_creation_input_tokens": 1_000_000, + "cache_creation": 12345, # malformed; must not raise + }, + ) + # Falls back to 5m default for the whole bucket. + assert _isclose( + out["cache_write_usd"], + 1_000_000 / 1_000_000.0 * base * ANTHROPIC_CACHE_5M_WRITE_MULT, + ) + + +def test_non_dict_server_tool_use_is_ignored(): + out = calculate_cost( + "anthropic", + "claude-opus-4-7", + {"input_tokens": 100, "output_tokens": 100, "server_tool_use": "garbage"}, + ) + assert out["server_tools_usd"] == 0.0 + + out = calculate_cost( + "openai", + "gpt-5.5", + {"input_tokens": 100, "output_tokens": 100, "openai_tool_use": [1, 2, 3]}, + ) + assert out["server_tools_usd"] == 0.0 + + +def test_non_dict_input_tokens_details_is_ignored(): + out = calculate_cost( + "openai", + "gpt-5.5", + { + "input_tokens": 100, + "output_tokens": 0, + "input_tokens_details": "nope", + "prompt_tokens_details": [1, 2, 3], + }, + ) + # No cached_tokens recovered -> no discount. + assert out["cache_read_usd"] == 0.0 + + +# ── unknown provider degrades gracefully ──────────────────────────── + + +def test_unknown_provider_priced_false_zero_bill(): + out = calculate_cost( + "gemini", + "gemini-pro", + {"input_tokens": 1_000_000, "output_tokens": 1_000_000}, + ) + assert out["priced"] is False + assert out["total_usd"] == 0.0 + # Tokens still report for the UI. + assert out["billable_input_tokens"] == 1_000_000 + assert out["billable_output_tokens"] == 1_000_000 + + +def test_anthropic_provider_with_openai_model_priced_false(): + # OpenAI id against Anthropic table must not falsely match. + out = calculate_cost( + "anthropic", + "gpt-5.5", + {"input_tokens": 1_000_000, "output_tokens": 0}, + ) + assert out["priced"] is False + + +# ── all-zero / empty usage stays at zero ──────────────────────────── + + +def test_empty_usage_dict_zero_bill(): + out = calculate_cost("openai", "gpt-5.5", {}) + assert out["priced"] is True # model is in the table + assert out["billable_input_tokens"] == 0 + assert out["total_usd"] == 0.0 + + +# ── Defense-in-depth: Anthropic prompt_tokens_details.cached_tokens ── + + +def test_anthropic_prompt_tokens_details_fallback_when_native_key_missing(): + """Chat-style envelope without `cache_read_input_tokens` but with + mirrored `prompt_tokens_details.cached_tokens` should still apply + the cache_read discount.""" + r = calculate_cost( + provider = "anthropic", + model = "claude-opus-4-7", + usage = { + "prompt_tokens": 1_000_000, + "completion_tokens": 0, + # Only the mirrored shape (no native key). + "prompt_tokens_details": {"cached_tokens": 1_000_000}, + "cache_creation_input_tokens": 0, + }, + ) + assert r["billable_input_tokens"] == 1_000_000, r + # 1M cached at 0.1x of $5 (opus 4.7) = $0.50 + assert math.isclose(r["cache_read_usd"], 0.5, rel_tol = 1e-3), r + + +def test_anthropic_native_key_takes_precedence_over_mirrored(): + """When both native and mirrored cache-read fields are present, + the native Anthropic field wins (mirror is fallback-only).""" + r = calculate_cost( + provider = "anthropic", + model = "claude-opus-4-7", + usage = { + "prompt_tokens": 1_000_000, + "cache_read_input_tokens": 800_000, + "prompt_tokens_details": {"cached_tokens": 1_000_000}, + "cache_creation_input_tokens": 0, + }, + ) + # billable = uncached_input + cache_creation + cache_read + # = (1M - 0 - 800k) + 0 + 800k = 1M + assert r["billable_input_tokens"] == 1_000_000, r + # cache_read uses 800k (native), not 1M (mirrored). + assert math.isclose(r["cache_read_usd"], 0.4, rel_tol = 1e-3), r + + +def test_anthropic_native_zero_takes_precedence_over_mirrored(): + """Explicit `cache_read_input_tokens: 0` is authoritative; a stale + mirrored block from a proxy must not inflate cache_read past it.""" + r = calculate_cost( + provider = "anthropic", + model = "claude-opus-4-7", + usage = { + "input_tokens": 1_000_000, + "output_tokens": 0, + "cache_read_input_tokens": 0, + # Stale mirror from a proxy; must be ignored (native present). + "prompt_tokens_details": {"cached_tokens": 1_000_000}, + }, + ) + # Native is 0 -> cache_read stays 0. + assert r["cache_read_usd"] == 0.0, r + # billable = input + cache_creation + cache_read = 1M + 0 + 0 + assert r["billable_input_tokens"] == 1_000_000, r + # 1M uncached at $5/M (no discount). + assert math.isclose(r["input_usd"], 5.0, rel_tol = 1e-3), r + assert math.isclose(r["total_usd"], 5.0, rel_tol = 1e-3), r + + +# ── _build_usage_chunk preserves cache_creation breakdown ── + + +def test_build_usage_chunk_forwards_anthropic_cache_creation_breakdown(): + """Chat-style envelope must carry the 5m/1h cache-write breakdown + so downstream cost calc applies the 2x 1h premium.""" + import json + from core.inference.external_provider import _build_usage_chunk + + chunk = _build_usage_chunk( + completion_id = "cmpl-x", + provider = "anthropic", + last_usage = { + "input_tokens": 10, + "output_tokens": 5, + "cache_creation_input_tokens": 1_000_000, + "cache_read_input_tokens": 0, + "cache_creation": { + "ephemeral_5m_input_tokens": 250_000, + "ephemeral_1h_input_tokens": 750_000, + }, + }, + ) + assert chunk is not None + payload = json.loads(chunk.split("data: ", 1)[1]) + cc = payload["usage"]["cache_creation"] + assert cc["ephemeral_1h_input_tokens"] == 750_000, cc + assert cc["ephemeral_5m_input_tokens"] == 250_000, cc + + +def test_calculate_cost_uses_forwarded_cache_creation_for_1h_premium(): + """Re-emitted chat envelope must price 1h cache writes at 2x base.""" + r = calculate_cost( + provider = "anthropic", + model = "claude-opus-4-7", + usage = { + "prompt_tokens": 1_000_010, + "completion_tokens": 0, + "cache_creation_input_tokens": 1_000_000, + "cache_read_input_tokens": 0, + "cache_creation": { + "ephemeral_5m_input_tokens": 0, + "ephemeral_1h_input_tokens": 1_000_000, + }, + }, + ) + # 1M at 1h-premium (2x of $5 = $10); 5m baseline would be $6.25. + assert math.isclose(r["cache_write_usd"], 10.0, rel_tol = 1e-2), r diff --git a/studio/frontend/src/components/assistant-ui/message-timing.tsx b/studio/frontend/src/components/assistant-ui/message-timing.tsx index 4fb68d4b90..31f742bc5e 100644 --- a/studio/frontend/src/components/assistant-ui/message-timing.tsx +++ b/studio/frontend/src/components/assistant-ui/message-timing.tsx @@ -33,10 +33,24 @@ export const MessageTiming: FC<{ if (timing?.totalStreamTime === undefined) return null; - const serverTimings = ( + const custom = ( message.metadata as Record | undefined - )?.custom as { serverTimings?: Record } | undefined; - const st = serverTimings?.serverTimings; + )?.custom as + | { + serverTimings?: Record; + contextUsage?: { + cachedTokens?: number; + cacheWriteTokens?: number; + }; + } + | undefined; + const st = custom?.serverTimings; + // `??` (not `||`) so an explicit cache_n=0 isn't replaced by a stale + // contextUsage.cachedTokens from a prior turn. + const cacheHits = + st?.cache_n ?? custom?.contextUsage?.cachedTokens ?? 0; + // Anthropic-only cache-write count. + const cacheWrites = custom?.contextUsage?.cacheWriteTokens ?? 0; // Guard unphysical tok/s: llama.cpp emits predicted_ms=0 on no-op // turns, blowing the rate up to Infinity. Require >=1 token AND a @@ -122,11 +136,19 @@ export const MessageTiming: FC<{ )} - {(st?.cache_n ?? 0) > 0 && ( + {cacheHits > 0 && (
Cache hits - {formatNumber(st!.cache_n)} + {formatNumber(cacheHits)} + +
+ )} + {cacheWrites > 0 && ( +
+ Cache writes + + {formatNumber(cacheWrites)}
)} @@ -146,7 +168,7 @@ export const MessageTiming: FC<{ ) : ( <> - {/* Client-side metrics (safetensors fallback) */} + {/* Client-side metrics (safetensors + external provider fallback) */} {timing.firstTokenTime !== undefined && (
First token @@ -155,6 +177,22 @@ export const MessageTiming: FC<{
)} + {cacheHits > 0 && ( +
+ Cache hits + + {formatNumber(cacheHits)} + +
+ )} + {cacheWrites > 0 && ( +
+ Cache writes + + {formatNumber(cacheWrites)} + +
+ )}
Total diff --git a/studio/frontend/src/components/assistant-ui/sources.tsx b/studio/frontend/src/components/assistant-ui/sources.tsx index 3a55c3fa78..140b61f932 100644 --- a/studio/frontend/src/components/assistant-ui/sources.tsx +++ b/studio/frontend/src/components/assistant-ui/sources.tsx @@ -127,6 +127,12 @@ function Source({ // ── Source badge with hover card ───────────────────────────── interface SourceData { + /** + * Stable per-citation key. Two Anthropic document citations into + * different spans of the same source share a ``url``, so React keys + * on ``id`` to keep each footnote distinct. + */ + id: string; url: string; title: string; description?: string; @@ -190,8 +196,14 @@ const SourcesGroup: FC = () => { "url" in part && part.url ) { + const url = part.url as string; + const partId = + typeof (part as { id?: unknown }).id === "string" + ? ((part as { id: string }).id) + : url; sources.push({ - url: part.url as string, + id: partId, + url, title: (part as { title?: string }).title || "", description: (part as { metadata?: { description?: string } }) .metadata?.description, @@ -258,7 +270,7 @@ const SourcesGroup: FC = () => { className="flex w-full flex-wrap gap-1 invisible absolute pointer-events-none" > {sources.map((source) => ( - + {source.title || extractDomain(source.url)} @@ -270,7 +282,7 @@ const SourcesGroup: FC = () => { {/* Visible container */}
{displayedSources.map((source) => ( - + ))} {shouldCollapse && !expanded && (
- {view.mode === "single" && ggufContextLength && contextUsage ? ( + {view.mode === "single" && contextUsage ? ( s.activeThreadId); const openAiApiKeyForSection = activeExternalProvider ? getExternalProviderApiKey(activeExternalProvider.id) || null @@ -1181,6 +1188,28 @@ export function ChatSettingsPanel({
) : null} + {showFastModeControl ? ( +
+
+ + Fast mode + + + Beta. Up to 2.5x higher output tokens per second on + Claude Opus 4.6 and 4.7 at 6x standard Opus pricing. + Switching between fast and standard invalidates the + prompt cache and is incompatible with the Priority + service tier. + +
+ +
+ ) : null} ) : null} diff --git a/studio/frontend/src/features/chat/components/context-usage-bar.tsx b/studio/frontend/src/features/chat/components/context-usage-bar.tsx index 91bdd35caa..e15a3e0b4f 100644 --- a/studio/frontend/src/features/chat/components/context-usage-bar.tsx +++ b/studio/frontend/src/features/chat/components/context-usage-bar.tsx @@ -28,37 +28,66 @@ function getSeverityColor(percent: number): { export const ContextUsageBar: FC<{ used: number; - total: number; + // null on external providers (no known window); bar then hides the ratio. + total?: number | null; cached?: number; + // Anthropic-only (billed at the write premium). + cacheWrites?: number; promptTokens?: number; completionTokens?: number; className?: string; -}> = ({ used, total, cached, promptTokens, completionTokens, className }) => { - if (total <= 0) return null; +}> = ({ + used, + total, + cached, + cacheWrites, + promptTokens, + completionTokens, + className, +}) => { + const hasKnownLimit = typeof total === "number" && total > 0; + const hasUsageDetails = + promptTokens !== undefined || + completionTokens !== undefined || + (cached !== undefined && cached > 0) || + (cacheWrites !== undefined && cacheWrites > 0); - const percent = Math.min((used / total) * 100, 100); - const severity = getSeverityColor(percent); + // Nothing to show: no limit and no per-turn counters. + if (!hasKnownLimit && used <= 0 && !hasUsageDetails) return null; + + const percent = hasKnownLimit + ? Math.min((used / (total as number)) * 100, 100) + : null; + const severity = getSeverityColor(percent ?? 0); return (
-
- Context usage - - {percent.toFixed(1)}% - -
+ {hasKnownLimit && percent !== null ? ( +
+ Context usage + + {percent.toFixed(1)}% + +
+ ) : null} {promptTokens !== undefined && (
Prompt tokens @@ -98,20 +129,32 @@ export const ContextUsageBar: FC<{
)} + {cacheWrites !== undefined && cacheWrites > 0 && ( +
+ Cache writes + + {formatTokenCountFull(cacheWrites)} + +
+ )}
- Total + + {hasKnownLimit ? "Total" : "Total tokens"} + - {formatTokenCountFull(used)} / {formatTokenCountFull(total)} + {hasKnownLimit + ? `${formatTokenCountFull(used)} / ${formatTokenCountFull(total as number)}` + : formatTokenCountFull(used)}
- {percent > 85 && ( + {hasKnownLimit && percent !== null && percent > 85 ? (
Close to the context limit. Generation will stop at 100%. Increase Context Length in the chat Settings panel to keep going.
- )} + ) : null}
diff --git a/studio/frontend/src/features/chat/provider-capabilities.ts b/studio/frontend/src/features/chat/provider-capabilities.ts index a04992d88e..e97d419b45 100644 --- a/studio/frontend/src/features/chat/provider-capabilities.ts +++ b/studio/frontend/src/features/chat/provider-capabilities.ts @@ -211,11 +211,9 @@ export function providerSupportsBuiltinWebSearch( /** * Whether the external provider exposes a server-side web_fetch tool - * that retrieves a single URL (text or PDF) and emits a document block. - * Only Anthropic ships one today (`web_fetch_20250910`); the chat - * composer pairs it with the Search pill because the typical workflow - * is "search returns URLs, fetch reads them" and the UI doesn't (yet) - * expose web_fetch as an independent toggle. + * (single URL, text or PDF) emitting a document block. Anthropic-only + * today (`web_fetch_20250910` / `web_fetch_20260209`). Gates the + * composer's standalone Fetch pill, independent of Search. */ export function providerSupportsBuiltinWebFetch( providerType: string | null | undefined, @@ -223,6 +221,30 @@ export function providerSupportsBuiltinWebFetch( return providerType === "anthropic"; } +/** + * Whether the active provider + model supports Anthropic fast-mode + * (`speed: "fast"` + `fast-mode-2026-02-01` header). Opus 4.6 / 4.7 + * only per https://platform.claude.com/docs/en/build-with-claude/fast-mode. + * Backend silently drops on unsupported models as a second defence. + */ +const ANTHROPIC_FAST_MODE_MODEL_PREFIXES = [ + "claude-opus-4-7", + "claude-opus-4-6", +] as const; + +export function providerSupportsFastMode( + providerType: string | null | undefined, + modelId: string | null | undefined, +): boolean { + if (providerType !== "anthropic") return false; + if (!modelId) return false; + // Family boundary ("" or "-") required so IDs like "claude-opus-4-70" + // / "claude-opus-4-7b" do not match. + return ANTHROPIC_FAST_MODE_MODEL_PREFIXES.some( + (prefix) => modelId === prefix || modelId.startsWith(`${prefix}-`), + ); +} + /** * Whether the selected external provider/model exposes a server-side * code-execution tool. Two providers ship one today: diff --git a/studio/frontend/src/features/chat/runtime-provider.tsx b/studio/frontend/src/features/chat/runtime-provider.tsx index d01383b309..21be5f6e3e 100644 --- a/studio/frontend/src/features/chat/runtime-provider.tsx +++ b/studio/frontend/src/features/chat/runtime-provider.tsx @@ -826,17 +826,24 @@ function useStudioRuntimeAdapters(): StudioRuntimeAdapters { completionTokens: number; totalTokens: number; cachedTokens: number; + cacheWriteTokens?: number; modelId?: string; } | undefined; const store = useChatRuntimeStore.getState(); - if ( - savedUsage && - store.ggufContextLength && - savedUsage.totalTokens <= store.ggufContextLength && - (!savedUsage.modelId || - savedUsage.modelId === store.params.checkpoint) - ) { + // Window check applies only when a local GGUF window is known; + // external providers have ggufContextLength === null. + const withinLocalLimit = + !store.ggufContextLength || + (savedUsage?.totalTokens ?? 0) <= store.ggufContextLength; + // Legacy unscoped usage (no modelId) is only trusted when a + // known local window bounds the totals, so we can't misattribute + // an old local turn to a newly-selected external provider. + const modelMatches = savedUsage?.modelId + ? savedUsage.modelId === store.params.checkpoint + : typeof store.ggufContextLength === "number" && + store.ggufContextLength > 0; + if (savedUsage && withinLocalLimit && modelMatches) { store.setContextUsage(savedUsage); } diff --git a/studio/frontend/src/features/chat/shared-composer.tsx b/studio/frontend/src/features/chat/shared-composer.tsx index 246fd81510..76e77d1288 100644 --- a/studio/frontend/src/features/chat/shared-composer.tsx +++ b/studio/frontend/src/features/chat/shared-composer.tsx @@ -21,7 +21,7 @@ import { isTauri } from "@/lib/api-base"; import { isMultimodalResponse } from "./types/api"; import { getImageInputUnavailableReason } from "./utils/image-input-support"; import { useAui } from "@assistant-ui/react"; -import { ArrowUpIcon, GlobeIcon, HeadphonesIcon, ImageIcon, LightbulbIcon, LightbulbOffIcon, MicIcon, PlusIcon, SquareIcon, XIcon } from "lucide-react"; +import { ArrowUpIcon, DownloadIcon, GlobeIcon, HeadphonesIcon, ImageIcon, LightbulbIcon, LightbulbOffIcon, MicIcon, PlusIcon, SquareIcon, XIcon } from "lucide-react"; import { toast } from "@/lib/toast"; import { loadModel, validateModel } from "./api/chat-api"; import { parseExternalModelId, providerTypeSupportsVision } from "./external-providers"; @@ -34,6 +34,7 @@ import { getExternalReasoningCapabilities, providerSupportsBuiltinCodeExecution, providerSupportsBuiltinImageGeneration, + providerSupportsBuiltinWebFetch, } from "./provider-capabilities"; import { type CompositionEvent, @@ -336,6 +337,12 @@ export function SharedComposer({ const setImageToolsEnabled = useChatRuntimeStore( (s) => s.setImageToolsEnabled, ); + const webFetchToolsEnabled = useChatRuntimeStore( + (s) => s.webFetchToolsEnabled, + ); + const setWebFetchToolsEnabled = useChatRuntimeStore( + (s) => s.setWebFetchToolsEnabled, + ); const lastOpenRouterChosenModel = useChatRuntimeStore( (s) => s.lastOpenRouterChosenModel, ); @@ -426,6 +433,9 @@ export function SharedComposer({ effectiveExternalModelId, selectedExternalProvider?.baseUrl, ); + const supportsBuiltinWebFetch = providerSupportsBuiltinWebFetch( + selectedExternalProvider?.providerType, + ); const searchDisabled = !modelLoaded || !(supportsTools || supportsBuiltinWebSearch); const codeDisabled = @@ -437,6 +447,9 @@ export function SharedComposer({ // the pill row stays compact for providers without the capability. const imageDisabled = !modelLoaded || !supportsBuiltinImageGeneration; const showImagePill = supportsBuiltinImageGeneration; + // Fetch pill: Anthropic-only (web_fetch_20250910 / web_fetch_20260209). + const webFetchDisabled = !modelLoaded || !supportsBuiltinWebFetch; + const showWebFetchPill = supportsBuiltinWebFetch; // Backwards-compatible alias for any other call site that may still // reference `toolsDisabled` (rare; both pills used it before). const toolsDisabled = codeDisabled; @@ -1106,6 +1119,23 @@ export function SharedComposer({ Images )} + {showWebFetchPill && ( + + )}
{dictationSupported && ( diff --git a/studio/frontend/src/features/chat/stores/chat-runtime-store.ts b/studio/frontend/src/features/chat/stores/chat-runtime-store.ts index dd26006228..7829d01f89 100644 --- a/studio/frontend/src/features/chat/stores/chat-runtime-store.ts +++ b/studio/frontend/src/features/chat/stores/chat-runtime-store.ts @@ -25,6 +25,8 @@ export const CHAT_REASONING_ENABLED_KEY = "unsloth_chat_reasoning_enabled"; export const CHAT_TOOLS_ENABLED_KEY = "unsloth_chat_tools_enabled"; export const CHAT_CODE_TOOLS_ENABLED_KEY = "unsloth_chat_code_tools_enabled"; export const CHAT_IMAGE_TOOLS_ENABLED_KEY = "unsloth_chat_image_tools_enabled"; +export const CHAT_WEB_FETCH_TOOLS_ENABLED_KEY = + "unsloth_chat_web_fetch_tools_enabled"; // External provider selection is encoded into `params.checkpoint` as // `external::::`. PersistedChatSettings deliberately @@ -262,9 +264,21 @@ type ChatRuntimeStore = { * receive the tool because their runtime cannot dispatch it. */ supportsBuiltinImageGeneration: boolean; + /** + * Whether the active external provider exposes a server-side + * web_fetch tool (Anthropic's `web_fetch_20250910` / + * `web_fetch_20260209`). Gates the composer's Fetch pill, + * independent of Search. + */ + supportsBuiltinWebFetch: boolean; toolsEnabled: boolean; codeToolsEnabled: boolean; imageToolsEnabled: boolean; + /** + * Fetch pill state, independent of `toolsEnabled` (Search). Only + * consulted when `providerSupportsBuiltinWebFetch` is true. + */ + webFetchToolsEnabled: boolean; toolStatus: string | null; generatingStatus: string | null; autoHealToolCalls: boolean; @@ -291,6 +305,8 @@ type ChatRuntimeStore = { completionTokens: number; totalTokens: number; cachedTokens: number; + // Anthropic-only; optional so pre-cache-stats persisted entries load. + cacheWriteTokens?: number; } | null; modelLoading: boolean; activeNativePathToken: string | null; @@ -324,6 +340,7 @@ type ChatRuntimeStore = { setToolsEnabled: (enabled: boolean, options?: { persist?: boolean }) => void; setCodeToolsEnabled: (enabled: boolean) => void; setImageToolsEnabled: (enabled: boolean) => void; + setWebFetchToolsEnabled: (enabled: boolean) => void; setToolStatus: (status: string | null) => void; setGeneratingStatus: (status: string | null) => void; setAutoHealToolCalls: (enabled: boolean) => void; @@ -382,6 +399,7 @@ const PERSISTED_INFERENCE_PARAM_KEYS = [ "maxTokens", "systemPrompt", "trustRemoteCode", + "fastMode", ] as const satisfies readonly PersistedInferenceParamKey[]; const SCALAR_SETTING_KEYS = [ @@ -569,9 +587,11 @@ export const useChatRuntimeStore = create((set, get) => ({ supportsBuiltinWebSearch: false, supportsBuiltinCodeExecution: false, supportsBuiltinImageGeneration: false, + supportsBuiltinWebFetch: false, toolsEnabled: loadBool(CHAT_TOOLS_ENABLED_KEY, false), codeToolsEnabled: loadBool(CHAT_CODE_TOOLS_ENABLED_KEY, false), imageToolsEnabled: loadBool(CHAT_IMAGE_TOOLS_ENABLED_KEY, false), + webFetchToolsEnabled: loadBool(CHAT_WEB_FETCH_TOOLS_ENABLED_KEY, false), toolStatus: null, generatingStatus: null, autoHealToolCalls: true, @@ -645,7 +665,14 @@ export const useChatRuntimeStore = create((set, get) => ({ if (state.settingsHydrated && hasKeys(changedParams)) { saveSettingsPatch({ inferenceParams: changedParams }); } - return { params }; + // Mirror setCheckpoint: the local model load path can mutate + // params.checkpoint via setParams() before setCheckpoint runs, + // leaving stale per-turn counters under the new checkpoint. + const checkpointChanged = state.params.checkpoint !== params.checkpoint; + return { + params, + ...(checkpointChanged ? { contextUsage: null } : {}), + }; }), setCustomPresets: (customPresets) => set(() => { @@ -709,12 +736,17 @@ export const useChatRuntimeStore = create((set, get) => ({ // mount, and a stale persisted local id would race against the // freshly-loaded model. See LAST_EXTERNAL_CHECKPOINT_KEY notes. saveLastExternalCheckpoint(isExternalModelId(modelId) ? modelId : null); + // Clear stale per-turn usage when the model changes; the relaxed + // external-provider render gate would otherwise show old counters + // until the next completion overwrites them. + const checkpointChanged = state.params.checkpoint !== modelId; return { params: { ...state.params, checkpoint: modelId, }, activeGgufVariant: ggufVariant ?? null, + ...(checkpointChanged ? { contextUsage: null } : {}), }; }), setActiveThreadId: (activeThreadId) => @@ -749,9 +781,11 @@ export const useChatRuntimeStore = create((set, get) => ({ supportsBuiltinWebSearch: false, supportsBuiltinCodeExecution: false, supportsBuiltinImageGeneration: false, + supportsBuiltinWebFetch: false, toolsEnabled: false, codeToolsEnabled: false, imageToolsEnabled: false, + webFetchToolsEnabled: false, toolStatus: null, kvCacheDtype: null, loadedKvCacheDtype: null, @@ -811,6 +845,11 @@ export const useChatRuntimeStore = create((set, get) => ({ saveBool(CHAT_IMAGE_TOOLS_ENABLED_KEY, imageToolsEnabled); return { imageToolsEnabled }; }), + setWebFetchToolsEnabled: (webFetchToolsEnabled) => + set(() => { + saveBool(CHAT_WEB_FETCH_TOOLS_ENABLED_KEY, webFetchToolsEnabled); + return { webFetchToolsEnabled }; + }), setToolStatus: (toolStatus) => set({ toolStatus }), setGeneratingStatus: (generatingStatus) => set({ generatingStatus }), setAutoHealToolCalls: (autoHealToolCalls) => diff --git a/studio/frontend/src/features/chat/types/api.ts b/studio/frontend/src/features/chat/types/api.ts index 7c5b5bf99f..3cf4d83a48 100644 --- a/studio/frontend/src/features/chat/types/api.ts +++ b/studio/frontend/src/features/chat/types/api.ts @@ -304,6 +304,12 @@ export interface OpenAIChatCompletionsRequest { * keeps each provider's upstream default. */ parallel_tool_calls?: boolean; + /** + * Anthropic fast-mode toggle. Opus 4.6 / 4.7 only; backend drops + * silently on every other model + provider. See + * https://platform.claude.com/docs/en/build-with-claude/fast-mode + */ + fast_mode?: boolean | null; } export interface OpenAIChatDelta { diff --git a/studio/frontend/src/features/chat/types/runtime.ts b/studio/frontend/src/features/chat/types/runtime.ts index 49c7c291e7..7fac74979b 100644 --- a/studio/frontend/src/features/chat/types/runtime.ts +++ b/studio/frontend/src/features/chat/types/runtime.ts @@ -50,6 +50,12 @@ export interface InferenceParams { checkpoint: string; /** Allow loading models with custom code (e.g. NVIDIA Nemotron). Only enable for repos you trust. */ trustRemoteCode?: boolean; + /** + * Anthropic fast-mode toggle. Opus 4.6 / 4.7 only; higher OTPS at + * 6x standard Opus pricing. Default false. + * https://platform.claude.com/docs/en/build-with-claude/fast-mode + */ + fastMode?: boolean; } export const DEFAULT_INFERENCE_PARAMS: InferenceParams = { @@ -69,6 +75,7 @@ export const DEFAULT_INFERENCE_PARAMS: InferenceParams = { systemPrompt: "", checkpoint: "", trustRemoteCode: false, + fastMode: false, }; export interface ChatModelSummary { diff --git a/studio/frontend/src/features/chat/utils/chat-settings-storage.ts b/studio/frontend/src/features/chat/utils/chat-settings-storage.ts index e6f6b693f1..1b186db727 100644 --- a/studio/frontend/src/features/chat/utils/chat-settings-storage.ts +++ b/studio/frontend/src/features/chat/utils/chat-settings-storage.ts @@ -182,6 +182,11 @@ function sanitizeInferenceParams( if (typeof value.parallelToolCalls === "boolean") { params.parallelToolCalls = value.parallelToolCalls; } + // Mirror trustRemoteCode handling so the toggle survives reload + // and the /api/chat/settings round-trip. + if (typeof value.fastMode === "boolean") { + params.fastMode = value.fastMode; + } return hasKeys(params) ? params : undefined; } diff --git a/tests/python/test_construct_chat_template_validation.py b/tests/python/test_construct_chat_template_validation.py new file mode 100644 index 0000000000..9ab68639c4 --- /dev/null +++ b/tests/python/test_construct_chat_template_validation.py @@ -0,0 +1,77 @@ +"""Negative-path validation tests for unsloth.chat_templates.construct_chat_template. + +Regression coverage for the str.find() / regex no-match guards added in +PR #5763 follow-up: missing placeholders or unrecoverable two-example +structures must raise RuntimeError with a clear message, not IndexError +or AttributeError, and must never silently drop the last character via +s[:-1]. + +Uses a minimal fake tokenizer so the cases run on CPU-only CI without +HF_TOKEN and without downloading a gated model. The validation paths +exercised here fail before construct_chat_template reaches any heavy +tokenizer interaction, so the stub stays small. +""" + +import pytest + +from unsloth.chat_templates import construct_chat_template + + +class _FakeTokenizer: + """Minimum surface construct_chat_template touches before the + validation guards fire.""" + + name_or_path = "fake/tokenizer" + eos_token = "" + + def get_vocab(self): + return {"": 0} + + +@pytest.mark.parametrize( + "template, expected_in_message", + [ + ("only {INPUT} here, no output marker", "{OUTPUT}"), + ("only {OUTPUT} here, no input marker", "{INPUT}"), + ("neither sentinel here, just literal text", "{INPUT}"), + ("neither sentinel here, just literal text", "{OUTPUT}"), + ], +) +def test_missing_placeholder_in_chat_template_raises(template, expected_in_message): + with pytest.raises(RuntimeError) as exc_info: + construct_chat_template( + tokenizer = _FakeTokenizer(), + chat_template = template, + extra_eos_tokens = [""], + ) + assert expected_in_message in str(exc_info.value) + + +def test_single_pair_template_raises_clear_error_not_attribute_error(): + """One {INPUT}/{OUTPUT} pair (rather than the required two) used to + crash with AttributeError on `found.group(1)` after the for-loop + broke without setting `found`. Must raise RuntimeError now.""" + template = "user: {INPUT}\nassistant: {OUTPUT}\n" + with pytest.raises(RuntimeError): + construct_chat_template( + tokenizer = _FakeTokenizer(), + chat_template = template, + extra_eos_tokens = [""], + ) + + +def test_error_message_excerpt_is_bounded(): + """Error messages must include a bounded excerpt of the offending + template, not dump arbitrarily large content into the traceback.""" + huge = ("garbage " * 5000) + "{INPUT}" # ~40 KB, missing {OUTPUT} + with pytest.raises(RuntimeError) as exc_info: + construct_chat_template( + tokenizer = _FakeTokenizer(), + chat_template = huge, + extra_eos_tokens = [""], + ) + msg = str(exc_info.value) + # Excerpt is repr-quoted and capped; total message should stay well + # under the template length. + assert len(msg) < 1000 + assert "{OUTPUT}" in msg diff --git a/unsloth/chat_templates.py b/unsloth/chat_templates.py index e8a34cbc60..956fcb2392 100644 --- a/unsloth/chat_templates.py +++ b/unsloth/chat_templates.py @@ -2461,17 +2461,40 @@ extra_eos_tokens = None, f"{left_changed}" ) except: - ending = chat_template[chat_template.find("{OUTPUT}") + len("{OUTPUT}"):] + output_pos = chat_template.find("{OUTPUT}") + input_pos = chat_template.find("{INPUT}") + if output_pos == -1 or input_pos == -1: + missing = [] + if input_pos == -1: missing.append("{INPUT}") + if output_pos == -1: missing.append("{OUTPUT}") + raise RuntimeError( + f"Unsloth: chat_template must contain {' and '.join(missing)} " + f"placeholder(s). Got: {chat_template[:200]!r}" + ) + ending = chat_template[output_pos + len("{OUTPUT}"):] ending = re.escape(ending) find_text = "{INPUT}" + ending + "(.+?{OUTPUT}" + ending + ")" response_part = re.findall(find_text, chat_template, flags = re.DOTALL | re.MULTILINE) + if len(response_part) == 0: + raise RuntimeError( + "Unsloth: Could not recover a two-example structure from chat_template. " + "Provide exactly two {INPUT}/{OUTPUT} pairs (and optionally {SYSTEM}). " + f"Got: {chat_template[:200]!r}" + ) response_part = response_part[0] + found = None for j in range(1, len(response_part)): try_find = re.escape(response_part[:j]) try: found = next(re.finditer("(" + try_find + ").+?\\{INPUT\\}", chat_template, flags = re.DOTALL | re.MULTILINE)) except: break + if found is None: + raise RuntimeError( + "Unsloth: Could not locate a separator between examples in chat_template. " + "Provide exactly two {INPUT}/{OUTPUT} pairs (and optionally {SYSTEM}). " + f"Got: {chat_template[:200]!r}" + ) separator = found.group(1) response_start = chat_template.find(response_part) @@ -2607,8 +2630,20 @@ extra_eos_tokens = None, jinja_template = "{{ bos_token }}" + jinja_template # Get instruction and output parts for train_on_inputs = False - input_part = input_part [:input_part .find("{INPUT}")] - output_part = output_part[:output_part.find("{OUTPUT}")] + input_idx = input_part .find("{INPUT}") + output_idx = output_part.find("{OUTPUT}") + if input_idx == -1: + raise RuntimeError( + f"Unsloth: The instruction section of the template must contain the " + f"'{{INPUT}}' placeholder. Section: {input_part[:200]!r}" + ) + if output_idx == -1: + raise RuntimeError( + f"Unsloth: The response section of the template must contain the " + f"'{{OUTPUT}}' placeholder. Section: {output_part[:200]!r}" + ) + input_part = input_part [:input_idx ] + output_part = output_part[:output_idx] return modelfile, jinja_template, input_part, output_part