Merge branch 'main' into feat/expose-sampling-params-core

Resolve 5 conflicts where main added Anthropic Opus 4.6/4.7 fast-mode
support that touches the same signatures PR 5711 extended:

  studio/backend/core/inference/external_provider.py
    Keep both PR 5711 sampling fields (frequency_penalty, seed, stop,
    service_tier, parallel_tool_calls) and main's fast_mode in the
    stream signature, docstring, and _stream_anthropic call site.
  studio/backend/models/inference.py
    Append fast_mode Field alongside PR 5711's new ChatCompletionRequest
    fields; both flow through the existing dispatch.
  studio/backend/routes/inference.py
    Forward fast_mode and the PR 5711 sampling fields to the stream
    generator together.
  studio/frontend/src/features/chat/types/api.ts
    Add fast_mode? to OpenAIChatCompletionsRequest after the PR 5711
    field block.
  studio/frontend/src/features/chat/utils/chat-settings-storage.ts
    Persist fastMode alongside seed / stop / serviceTier /
    parallelToolCalls.

No semantic changes to either feature surface. 161 backend routing
tests still pass; frontend tsc clean.
This commit is contained in:
Daniel Han 2026-05-26 07:22:06 +00:00
commit c43c48a7e0
35 changed files with 4846 additions and 335 deletions

View file

@ -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:

View file

@ -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,

View file

@ -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("</think>")
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<sid>``
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<sid>`` 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<sid>`` 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("</think>")
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("</think>")
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("</think>")
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("</think>")
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("</think>")
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("</think>")
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("</think>")
reasoning_open = False
yield _chunk_with_text(tail_flushed)
if reasoning_open:
yield _chunk_with_text("</think>")
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("</think>")
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("</think>")
reasoning_open = False
yield _chunk_with_text(tail_flushed)
if reasoning_open:
yield _chunk_with_text("</think>")
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

View file

@ -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,
},

View file

@ -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

View file

@ -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":

View file

@ -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:

View file

@ -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

View file

@ -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("")

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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,<script>alert(1)</script>"
)
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

View file

@ -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},
}

View file

@ -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

View file

@ -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<sid>[\\ue202<sid>...][\\ue202<loc>]\\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<sid>`` 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<sid>`` 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

View file

@ -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

View file

@ -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

View file

@ -33,10 +33,24 @@ export const MessageTiming: FC<{
if (timing?.totalStreamTime === undefined) return null;
const serverTimings = (
const custom = (
message.metadata as Record<string, unknown> | undefined
)?.custom as { serverTimings?: Record<string, number> } | undefined;
const st = serverTimings?.serverTimings;
)?.custom as
| {
serverTimings?: Record<string, number>;
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<{
</span>
</div>
)}
{(st?.cache_n ?? 0) > 0 && (
{cacheHits > 0 && (
<div className="flex items-center justify-between gap-4">
<span className="text-muted-foreground">Cache hits</span>
<span className="font-mono tabular-nums">
{formatNumber(st!.cache_n)}
{formatNumber(cacheHits)}
</span>
</div>
)}
{cacheWrites > 0 && (
<div className="flex items-center justify-between gap-4">
<span className="text-muted-foreground">Cache writes</span>
<span className="font-mono tabular-nums">
{formatNumber(cacheWrites)}
</span>
</div>
)}
@ -146,7 +168,7 @@ export const MessageTiming: FC<{
</>
) : (
<>
{/* Client-side metrics (safetensors fallback) */}
{/* Client-side metrics (safetensors + external provider fallback) */}
{timing.firstTokenTime !== undefined && (
<div className="flex items-center justify-between gap-4">
<span className="text-muted-foreground">First token</span>
@ -155,6 +177,22 @@ export const MessageTiming: FC<{
</span>
</div>
)}
{cacheHits > 0 && (
<div className="flex items-center justify-between gap-4">
<span className="text-muted-foreground">Cache hits</span>
<span className="font-mono tabular-nums">
{formatNumber(cacheHits)}
</span>
</div>
)}
{cacheWrites > 0 && (
<div className="flex items-center justify-between gap-4">
<span className="text-muted-foreground">Cache writes</span>
<span className="font-mono tabular-nums">
{formatNumber(cacheWrites)}
</span>
</div>
)}
<div className="flex items-center justify-between gap-4">
<span className="text-muted-foreground">Total</span>
<span className="font-mono tabular-nums">

View file

@ -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) => (
<span key={source.url} className="inline-block">
<span key={source.id} className="inline-block">
<Source href={source.url}>
<SourceIcon url={source.url} />
<SourceTitle>{source.title || extractDomain(source.url)}</SourceTitle>
@ -270,7 +282,7 @@ const SourcesGroup: FC = () => {
{/* Visible container */}
<div className="flex flex-wrap gap-1">
{displayedSources.map((source) => (
<SourceBadge key={source.url} source={source} />
<SourceBadge key={source.id} source={source} />
))}
{shouldCollapse && !expanded && (
<button

View file

@ -24,6 +24,22 @@ const RE_TITLE = /Title:\s*(.+)/;
const RE_URL = /URL:\s*(.+)/;
const RE_SNIPPET = /Snippet:\s*(.+)/s;
/**
* Reject anything that is not a real http(s) URL. Web-search / web-fetch
* output is provider-controlled, so hostile ``javascript:`` / ``data:``
* lines must not reach the Source badge's <a href>.
*/
function isSafeHttpUrl(raw: string): boolean {
const value = raw.trim();
if (!value || /[\r\n]/.test(value)) return false;
try {
const parsed = new URL(value);
return parsed.protocol === "http:" || parsed.protocol === "https:";
} catch {
return false;
}
}
/** Parse the backend's "Title: ...\nURL: ...\nSnippet: ...\n---" format into structured sources. */
function parseSearchResults(raw: string): ParsedSource[] {
if (!raw) {
@ -35,13 +51,14 @@ function parseSearchResults(raw: string): ParsedSource[] {
const titleMatch = block.match(RE_TITLE);
const urlMatch = block.match(RE_URL);
const snippetMatch = block.match(RE_SNIPPET);
if (titleMatch && urlMatch) {
sources.push({
title: titleMatch[1].trim(),
url: urlMatch[1].trim(),
snippet: snippetMatch?.[1]?.trim() ?? "",
});
}
if (!titleMatch || !urlMatch) continue;
const url = urlMatch[1].trim();
if (!isSafeHttpUrl(url)) continue;
sources.push({
title: titleMatch[1].trim(),
url,
snippet: snippetMatch?.[1]?.trim() ?? "",
});
}
return sources;
}

View file

@ -28,6 +28,7 @@ import {
providerSupportsBuiltinImageGeneration,
providerSupportsBuiltinWebFetch,
providerSupportsBuiltinWebSearch,
providerSupportsFastMode,
} from "../provider-capabilities";
import { useChatRuntimeStore } from "../stores/chat-runtime-store";
import { useExternalProvidersStore } from "../stores/external-providers-store";
@ -70,6 +71,13 @@ interface ServerUsage {
prompt_tokens: number;
completion_tokens: number;
total_tokens: number;
// External prompt-cache fields (see _build_usage_chunk in
// external_provider.py). cache_creation is Anthropic-only.
prompt_tokens_details?: {
cached_tokens?: number;
};
cache_creation_input_tokens?: number;
cache_read_input_tokens?: number;
}
/** Server-side timing data from llama-server's timings object. */
@ -143,6 +151,91 @@ async function updateStoredChatThreadEventually(
}
}
/**
* Return ``raw`` when it is a safe-to-navigate http(s) URL, or "" otherwise.
* Rejects non-string input, CR/LF (header injection), and non-http(s)
* schemes (``javascript:`` / ``data:`` / ``vbscript:``) so provider /
* tool-controlled strings cannot land in an <a href>.
*/
function isSafeNavigableSourceUrl(raw: unknown): string {
if (typeof raw !== "string") return "";
const value = raw.trim();
if (!value || /[\r\n]/.test(value)) return "";
try {
const parsed = new URL(value);
if (parsed.protocol === "http:" || parsed.protocol === "https:") {
return value;
}
} catch {
// Fall through.
}
return "";
}
/** Convert an Anthropic document citation dict into a Sources-panel source. */
function documentCitationToSource(
cit: Record<string, unknown>,
fallbackIdx: number,
): {
type: "source";
sourceType: "url";
id: string;
url: string;
title: string;
metadata?: { description: string };
} | null {
const source =
typeof cit.source === "string" && cit.source ? cit.source : "";
const docTitle =
(typeof cit.document_title === "string" && cit.document_title) ||
(typeof cit.title === "string" && cit.title) ||
"";
const docIndex =
typeof cit.document_index === "number" ? cit.document_index : undefined;
// Only treat ``source`` as a navigable URL when it is real http(s);
// search_result_location can carry a free-form id (e.g. ``kb-doc-42``)
// or a hostile ``javascript:`` / ``data:`` / ``vbscript:`` string.
// Fall back to a stable doc anchor otherwise.
const url =
isSafeNavigableSourceUrl(source) || `#anthropic-doc-${docIndex ?? fallbackIdx}`;
const title = docTitle || source || `Document ${fallbackIdx + 1}`;
const cited =
typeof cit.cited_text === "string" ? cit.cited_text.trim() : "";
// Trim the cited snippet so the Sources panel stays scannable.
const description =
cited.length > 240 ? `${cited.slice(0, 240)}...` : cited;
// Anthropic numbers inline [N] per citation, not per source URL.
// Fold citation type + position-bearing fields into the id so two
// distinct citations on the same source (or two search_result_locations
// with different search_result_index) keep separate Sources entries.
const citationType =
typeof cit.type === "string" ? String(cit.type) : "";
const positionParts = [
cit.search_result_index,
cit.start_char_index,
cit.end_char_index,
cit.start_page_number,
cit.end_page_number,
cit.start_block_index,
cit.end_block_index,
]
.filter((v) => typeof v === "number")
.map((v) => String(v))
.join(":");
const idAnchor = positionParts
? `${citationType}:${positionParts}`
: `${citationType}:${fallbackIdx}`;
const id = `${url}#${idAnchor}`;
return {
type: "source" as const,
sourceType: "url" as const,
id,
url,
title,
...(description ? { metadata: { description } } : {}),
};
}
/** Parse "Title: ...\nURL: ...\nSnippet: ..." blocks into source content parts. */
function parseSourcesFromResult(raw: string): {
type: "source";
@ -167,7 +260,11 @@ function parseSourcesFromResult(raw: string): {
const urlMatch = block.match(/URL:\s*(.+)/);
const snippetMatch = block.match(/Snippet:\s*(.+)/);
if (titleMatch && urlMatch) {
const url = urlMatch[1].trim();
// Drop blocks whose ``URL:`` is not safe http(s); provider/tool
// output is attacker-controllable so a hostile ``javascript:`` /
// ``data:`` line must not reach the Sources panel <a href>.
const url = isSafeNavigableSourceUrl(urlMatch[1]);
if (!url) continue;
const snippet = snippetMatch?.[1]?.trim();
sources.push({
type: "source" as const,
@ -340,6 +437,18 @@ function collectImageParts(
return parts;
}
// Refusal flag stamped on assistant metadata when the backend emits the
// `anthropic_refusal` _toolEvent. We drop the refused pair from the next
// request body (Anthropic guidance: leaving refusals in context keeps
// refusing). Metadata (not text) prevents content from spoofing a reset.
function isAnthropicRefusalMessage(message: RunMessage): boolean {
if (message.role !== "assistant") return false;
const metadata = (message as { metadata?: unknown }).metadata as
| { custom?: Record<string, unknown> }
| undefined;
return metadata?.custom?.anthropicRefusal === true;
}
function toOpenAIMessage(message: RunMessage): {
role: "system" | "user" | "assistant";
content: OpenAIMessageContent;
@ -360,6 +469,11 @@ function toOpenAIMessage(message: RunMessage): {
/data:audio\/[a-z0-9.+-]+;base64,[A-Za-z0-9+/=]+/g,
"[audio]",
);
if (isAnthropicRefusalMessage(message)) {
// Prune refused assistant turn from outbound history; the
// rendered transcript still shows the user-visible notice.
return null;
}
}
const imageParts = collectImageParts(message);
@ -833,7 +947,13 @@ export function createOpenAIStreamAdapter(): ChatModelAdapter {
// Re-read store after potential auto-load / model ready wait
runtime = useChatRuntimeStore.getState();
const { params } = runtime;
const { supportsTools, toolsEnabled, codeToolsEnabled, imageToolsEnabled } = runtime;
const {
supportsTools,
toolsEnabled,
codeToolsEnabled,
imageToolsEnabled,
webFetchToolsEnabled,
} = runtime;
const externalSelection = parseExternalModelId(params.checkpoint);
const isExternalRequest = externalSelection !== null;
if (
@ -889,14 +1009,14 @@ export function createOpenAIStreamAdapter(): ChatModelAdapter {
externalProvider.baseUrl,
),
);
// web_fetch shares the Search pill with web_search (no separate
// UI toggle), so it follows toolsEnabled. Anthropic is the only
// provider that ships it today; on others providerSupportsBuiltinWebFetch
// returns false and this stays inert.
// Fetch pill is independent of Search (Anthropic bills web_fetch
// separately from web_search). Sourced from `webFetchToolsEnabled`;
// on providers without web_fetch the toggle is forced off in
// chat-page's runtime setState.
const webFetchEnabledForThisTurn =
Boolean(
externalProvider &&
toolsEnabled &&
webFetchToolsEnabled &&
providerSupportsBuiltinWebFetch(externalProvider.providerType),
);
const providerShipsWebFetch = Boolean(
@ -918,7 +1038,24 @@ export function createOpenAIStreamAdapter(): ChatModelAdapter {
),
);
const outboundMessages = messages
// Two-pass build: a refused assistant turn also drops the user
// prompt that triggered it (leaving it in context re-triggers
// the classifier). Refusal flag rides assistant
// metadata.custom.anthropicRefusal, set out-of-band from the
// backend _toolEvent.
const survivingMessages: RunMessage[] = [];
for (const message of messages) {
if (isAnthropicRefusalMessage(message)) {
const last = survivingMessages.at(-1);
if (last && last.role === "user") {
survivingMessages.pop();
}
continue;
}
survivingMessages.push(message);
}
const outboundMessages = survivingMessages
.map(toOpenAIMessage)
.filter((message): message is NonNullable<typeof message> =>
Boolean(message),
@ -941,14 +1078,20 @@ export function createOpenAIStreamAdapter(): ChatModelAdapter {
const webLabel = providerShipsWebFetch
? "web search or web fetch"
: "web search";
if (!webSearchEnabledForThisTurn && !codeExecEnabledForThisTurn) {
// Treat search and fetch as a single "any web tool" axis so
// the guard only warns when neither pill is on; checking
// webSearchEnabledForThisTurn alone mis-fired when only Fetch
// was on and suppressed live web_fetch calls.
const anyWebEnabledForThisTurn =
webSearchEnabledForThisTurn || webFetchEnabledForThisTurn;
if (!anyWebEnabledForThisTurn && !codeExecEnabledForThisTurn) {
disabledToolGuard =
`You do not have ${webLabel} or code execution tools in this conversation. ` +
"Answer from your own knowledge. " +
"If a request genuinely requires tool use, live data fetch or running code, " +
"inform the user that you do not have access to these capabilities. " +
"Do not return tool-call syntax inside your response.";
} else if (!webSearchEnabledForThisTurn) {
} else if (!anyWebEnabledForThisTurn) {
disabledToolGuard =
`You do not have ${webLabel} tools in this conversation. ` +
"You may still use code execution tools when they are available and useful. " +
@ -988,8 +1131,10 @@ export function createOpenAIStreamAdapter(): ChatModelAdapter {
});
}
}
const imageBase64 = findLatestUserImageBase64(messages);
const audioBase64 = findLatestUserAudioBase64(messages);
// Scan post-prune history so a refused user turn's image/audio
// doesn't gate or mis-attribute the next non-refused turn.
const imageBase64 = findLatestUserImageBase64(survivingMessages);
const audioBase64 = findLatestUserAudioBase64(survivingMessages);
// Block when ANY image is in the outbound payload (current or
// prior turns) and the loaded model can't process images. Keeps
@ -1025,7 +1170,7 @@ export function createOpenAIStreamAdapter(): ChatModelAdapter {
if (audioBase64) {
const audioName = runtime.pendingAudioName;
if (audioName) {
const lastUserMsg = [...messages]
const lastUserMsg = [...survivingMessages]
.reverse()
.find((m) => m.role === "user");
if (lastUserMsg) sentAudioNames.set(lastUserMsg.id, audioName);
@ -1136,6 +1281,21 @@ export function createOpenAIStreamAdapter(): ChatModelAdapter {
// Tool call content parts — accumulated and yielded cumulatively.
// result is set directly on the tool-call part when tool_end arrives.
const toolCallParts: ToolCallMessagePart[] = [];
// Anthropic document_citations tool_event payload, converted to
// Sources-panel source parts at end-of-stream so the inline [N]
// markers have matching entries.
const documentCitationParts: Array<{
type: "source";
sourceType: "url";
id: string;
url: string;
title: string;
metadata?: { description: string };
}> = [];
// Latched on the `anthropic_refusal` tool event; stamped onto the
// final assistant metadata as `custom.anthropicRefusal` to drive
// the history-prune above.
let anthropicRefusalSeen = false;
let serverMetadata: {
usage?: ServerUsage;
timings?: ServerTimings;
@ -1443,13 +1603,8 @@ export function createOpenAIStreamAdapter(): ChatModelAdapter {
enable_tools: true,
enabled_tools: [
...(webSearchEnabledForThisTurn ? ["web_search"] : []),
// Pair web_fetch with the Search pill on any
// provider that ships it (Anthropic today). The
// common workflow is "search returns URLs, fetch
// reads them"; without web_fetch the model can
// surface a citation but cannot quote from the
// page body, which is the whole point of the
// tool. There is no separate UI toggle yet.
// web_fetch has its own Fetch pill, independent
// of Search. Anthropic-only today.
...(webFetchEnabledForThisTurn ? ["web_fetch"] : []),
...(codeExecEnabledForThisTurn ? ["code_execution"] : []),
// OpenAI Responses-API only: `image_generation`
@ -1502,6 +1657,16 @@ export function createOpenAIStreamAdapter(): ChatModelAdapter {
isPromptCacheTtl(externalProvider.promptCacheTtl)
? { prompt_cache_ttl: externalProvider.promptCacheTtl }
: {}),
// Anthropic fast mode (Opus 4.6 / 4.7 only); backend
// silently drops on unsupported models as a second
// line of defence.
...(params.fastMode &&
providerSupportsFastMode(
externalProvider.providerType,
externalSelection.modelId,
)
? { fast_mode: true }
: {}),
...(externalReasoningCaps.supportsReasoning
? externalReasoningCaps.reasoningStyle === "reasoning_effort"
? externalReasoningEnabled
@ -1623,6 +1788,27 @@ export function createOpenAIStreamAdapter(): ChatModelAdapter {
}
continue;
}
if (toolEvent.type === "document_citations") {
// Convert Anthropic citations_delta footnotes into
// Sources-panel entries matching the inline [N] markers.
const cits = toolEvent.citations;
if (Array.isArray(cits)) {
cits.forEach((entry, idx) => {
if (!entry || typeof entry !== "object") return;
const part = documentCitationToSource(
entry as Record<string, unknown>,
idx,
);
if (
part &&
!documentCitationParts.some((p) => p.id === part.id)
) {
documentCitationParts.push(part);
}
});
}
continue;
}
if (toolEvent.type === "container_invalidated") {
if (resolvedThreadId) {
const field =
@ -1636,6 +1822,12 @@ export function createOpenAIStreamAdapter(): ChatModelAdapter {
}
continue;
}
if (toolEvent.type === "anthropic_refusal") {
// Latch the backend refusal signal so the final
// message metadata can drive the prune.
anthropicRefusalSeen = true;
continue;
}
if (toolEvent.type === "tool_start") {
const id =
(toolEvent.tool_call_id as string) ||
@ -1921,18 +2113,31 @@ export function createOpenAIStreamAdapter(): ChatModelAdapter {
const finalTokPerSec = meta?.timings?.predicted_per_second;
const serverPromptEvalTime = meta?.timings?.prompt_ms;
// Update context usage in store if we got valid server data
// Prefer llama-server timings; fall back to provider usage envelope.
const cachedTokens =
meta?.timings?.cache_n ??
meta?.usage?.prompt_tokens_details?.cached_tokens ??
meta?.usage?.cache_read_input_tokens ??
0;
// Anthropic-only (billed at the write premium).
const cacheWriteTokens = meta?.usage?.cache_creation_input_tokens ?? 0;
// Gate on the captured checkpoint still being active so a late
// completion from provider A doesn't populate the bar after the
// user switched to provider B mid-stream.
if (
meta?.usage &&
typeof meta.usage.prompt_tokens === "number" &&
typeof meta.usage.completion_tokens === "number" &&
typeof meta.usage.total_tokens === "number"
typeof meta.usage.total_tokens === "number" &&
useChatRuntimeStore.getState().params.checkpoint === params.checkpoint
) {
useChatRuntimeStore.getState().setContextUsage({
promptTokens: meta.usage.prompt_tokens,
completionTokens: meta.usage.completion_tokens,
totalTokens: meta.usage.total_tokens,
cachedTokens: meta.timings?.cache_n ?? 0,
cachedTokens,
cacheWriteTokens,
});
}
@ -1951,18 +2156,22 @@ export function createOpenAIStreamAdapter(): ChatModelAdapter {
...toolCallParts,
...parseAssistantContent(cumulativeText),
...sourceParts,
...documentCitationParts,
],
metadata: {
timing: finalTiming,
custom: {
reasoningDuration,
// Persisted refusal flag driving the two-pass prune.
anthropicRefusal: anthropicRefusalSeen || undefined,
serverTimings: meta?.timings ?? undefined,
contextUsage: meta?.usage
? {
promptTokens: meta.usage.prompt_tokens,
completionTokens: meta.usage.completion_tokens,
totalTokens: meta.usage.total_tokens,
cachedTokens: meta.timings?.cache_n ?? 0,
cachedTokens,
cacheWriteTokens,
modelId: params.checkpoint,
}
: undefined,

View file

@ -57,6 +57,7 @@ import {
getProviderCapabilities,
providerSupportsBuiltinCodeExecution,
providerSupportsBuiltinImageGeneration,
providerSupportsBuiltinWebFetch,
providerSupportsBuiltinWebSearch,
} from "./provider-capabilities";
import { ChatRuntimeProvider } from "./runtime-provider";
@ -71,6 +72,7 @@ import {
CHAT_CODE_TOOLS_ENABLED_KEY,
CHAT_IMAGE_TOOLS_ENABLED_KEY,
CHAT_TOOLS_ENABLED_KEY,
CHAT_WEB_FETCH_TOOLS_ENABLED_KEY,
loadOptionalBool,
useChatRuntimeStore,
} from "./stores/chat-runtime-store";
@ -782,6 +784,9 @@ export function ChatPage(): ReactElement {
selection.modelId,
provider?.baseUrl,
);
const supportsBuiltinWebFetch = providerSupportsBuiltinWebFetch(
provider?.providerType,
);
// Kimi's k2.6/k2.5 default to thinking enabled on the server side
// (per https://platform.kimi.ai/docs/models). Mirror that default
// in the UI so the Think pill comes up clicked when the user picks
@ -804,6 +809,9 @@ export function ChatPage(): ReactElement {
const storedImageToolsEnabled = loadOptionalBool(
CHAT_IMAGE_TOOLS_ENABLED_KEY,
);
const storedWebFetchToolsEnabled = loadOptionalBool(
CHAT_WEB_FETCH_TOOLS_ENABLED_KEY,
);
const nextToolsEnabled = supportsBuiltinWebSearch
? isKimi
? false
@ -837,6 +845,7 @@ export function ChatPage(): ReactElement {
supportsBuiltinWebSearch,
supportsBuiltinCodeExecution,
supportsBuiltinImageGeneration,
supportsBuiltinWebFetch,
toolsEnabled: nextToolsEnabled,
codeToolsEnabled: supportsBuiltinCodeExecution
? (storedCodeToolsEnabled ?? false)
@ -844,6 +853,10 @@ export function ChatPage(): ReactElement {
imageToolsEnabled: supportsBuiltinImageGeneration
? (storedImageToolsEnabled ?? false)
: false,
// Default Fetch off (Anthropic bills per fetch); deliberate opt-in.
webFetchToolsEnabled: supportsBuiltinWebFetch
? (storedWebFetchToolsEnabled ?? false)
: false,
});
}, [externalProvidersForChat, inferenceParams.checkpoint]);
const canCompare = useMemo(() => {
@ -1011,6 +1024,9 @@ export function ChatPage(): ReactElement {
selectedExternal?.modelId,
selectedProvider?.baseUrl,
);
const supportsBuiltinWebFetch = providerSupportsBuiltinWebFetch(
selectedProvider?.providerType,
);
// See sibling useEffect above: Kimi's k2.x default to thinking
// enabled, so the Think pill comes up clicked. Search pill stays
// off by default; mutual exclusion flips them via the composer.
@ -1029,6 +1045,9 @@ export function ChatPage(): ReactElement {
const storedImageToolsEnabled = loadOptionalBool(
CHAT_IMAGE_TOOLS_ENABLED_KEY,
);
const storedWebFetchToolsEnabled = loadOptionalBool(
CHAT_WEB_FETCH_TOOLS_ENABLED_KEY,
);
const nextToolsEnabled = supportsBuiltinWebSearch
? isKimi
? false
@ -1040,6 +1059,10 @@ export function ChatPage(): ReactElement {
ggufMaxContextLength: null,
ggufNativeContextLength: null,
activeNativePathToken: null,
// Clear previous-model counters; the relaxed external-provider
// render gate would otherwise show stale stats until the next
// completion overwrites them.
contextUsage: null,
supportsReasoning: reasoningCaps.supportsReasoning,
reasoningAlwaysOn: reasoningCaps.reasoningAlwaysOn,
reasoningStyle: reasoningCaps.reasoningStyle,
@ -1066,6 +1089,7 @@ export function ChatPage(): ReactElement {
supportsBuiltinWebSearch,
supportsBuiltinCodeExecution,
supportsBuiltinImageGeneration,
supportsBuiltinWebFetch,
toolsEnabled: nextToolsEnabled,
codeToolsEnabled: supportsBuiltinCodeExecution
? (storedCodeToolsEnabled ?? false)
@ -1073,6 +1097,9 @@ export function ChatPage(): ReactElement {
imageToolsEnabled: supportsBuiltinImageGeneration
? (storedImageToolsEnabled ?? false)
: false,
webFetchToolsEnabled: supportsBuiltinWebFetch
? (storedWebFetchToolsEnabled ?? false)
: false,
...(stillOnOpenRouterFree ? {} : { lastOpenRouterChosenModel: null }),
});
return;
@ -1164,7 +1191,9 @@ export function ChatPage(): ReactElement {
if (!saved) return;
viewBeforeCompareRef.current = null;
navigate({ to: "/chat", search: saved });
// Restore context usage from the active thread's last assistant message.
// Restore usage from the last assistant message, but only if it
// matches the currently active checkpoint. Without this guard the
// relaxed render gate would show stale stats from another model.
const threadId =
saved.thread ?? useChatRuntimeStore.getState().activeThreadId;
if (threadId) {
@ -1178,7 +1207,29 @@ export function ChatPage(): ReactElement {
const usage = metadata?.contextUsage as ReturnType<
typeof useChatRuntimeStore.getState
>["contextUsage"];
if (usage) useChatRuntimeStore.getState().setContextUsage(usage);
if (!usage) return;
const store = useChatRuntimeStore.getState();
const activeCheckpoint = store.params.checkpoint;
const usageModelId =
(usage as { modelId?: unknown }).modelId;
// Scope by modelId when present; reject if no active checkpoint
// (model-scoped usage cannot be attributed to "nothing").
if (typeof usageModelId === "string" && usageModelId) {
if (!activeCheckpoint || usageModelId !== activeCheckpoint) {
return;
}
}
// For local turns, also require the restored count to fit in
// the active window. Skip when unknown (external provider).
const limit = store.ggufContextLength;
if (
typeof limit === "number" &&
limit > 0 &&
(usage.totalTokens ?? 0) > limit
) {
return;
}
store.setContextUsage(usage);
})
.catch((error) => {
if (!isExpectedBackgroundChatStorageError(error)) {
@ -1494,11 +1545,13 @@ export function ChatPage(): ReactElement {
) : null}
</div>
<div className="ml-auto flex items-center gap-2">
{view.mode === "single" && ggufContextLength && contextUsage ? (
{view.mode === "single" && contextUsage ? (
<ContextUsageBar
used={contextUsage.totalTokens}
// null on external providers; the bar handles that.
total={ggufContextLength}
cached={contextUsage.cachedTokens}
cacheWrites={contextUsage.cacheWriteTokens}
promptTokens={contextUsage.promptTokens}
completionTokens={contextUsage.completionTokens}
className="h-[34px]"

View file

@ -89,6 +89,7 @@ import {
getProviderStopMax,
getServiceTierOptions,
providerSupportsBuiltinCodeExecution,
providerSupportsFastMode,
} from "./provider-capabilities";
import { StopSequencesInput } from "@/components/ui/stop-sequences-input";
import { useChatRuntimeStore } from "./stores/chat-runtime-store";
@ -581,6 +582,12 @@ export function ChatSettingsPanel({
activeExternalProvider.baseUrl,
) &&
activeExternalProvider.providerType === "openai";
const showFastModeControl =
activeExternalProvider != null &&
providerSupportsFastMode(
activeExternalProvider.providerType,
externalSelection?.modelId,
);
const activeThreadId = useChatRuntimeStore((s) => s.activeThreadId);
const openAiApiKeyForSection = activeExternalProvider
? getExternalProviderApiKey(activeExternalProvider.id) || null
@ -1181,6 +1188,28 @@ export function ChatSettingsPanel({
</Select>
</div>
) : null}
{showFastModeControl ? (
<div className="flex items-center justify-between gap-3 pt-3">
<div className="flex min-w-0 items-center gap-1.5">
<span className="min-w-0 text-[13px] font-medium leading-[1.25] tracking-nav text-nav-fg">
Fast mode
</span>
<InfoHint>
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.
</InfoHint>
</div>
<Switch
className="panel-switch shrink-0"
checked={Boolean(params.fastMode)}
onCheckedChange={set("fastMode")}
aria-label="Fast mode"
/>
</div>
) : null}
</CollapsibleSection>
) : null}

View file

@ -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 (
<Tooltip>
<TooltipTrigger asChild>
<button
type="button"
aria-label={`Context usage: ${formatTokenCount(used)} of ${formatTokenCount(total)} tokens`}
aria-label={
hasKnownLimit
? `Context usage: ${formatTokenCount(used)} of ${formatTokenCount(total as number)} tokens`
: `Token usage: ${formatTokenCount(used)} tokens`
}
className={cn(
"flex items-center gap-2 rounded-[10px] px-2.5 py-1 font-mono text-chat-icon-fg text-[13px] tabular-nums transition-colors hover:bg-chat-icon-bg-hover hover:text-chat-icon-fg-hover",
className,
)}
>
<span>
{formatTokenCount(used)} / {formatTokenCount(total)}
{hasKnownLimit
? `${formatTokenCount(used)} / ${formatTokenCount(total as number)}`
: `${formatTokenCount(used)} tokens`}
</span>
<div className="h-1.5 w-16 rounded-full bg-black/10 dark:bg-white/15 overflow-hidden">
<div
className={cn("h-full rounded-full transition-all", severity.bar)}
style={{ width: `${percent}%` }}
/>
</div>
{hasKnownLimit && percent !== null ? (
<div className="h-1.5 w-16 rounded-full bg-black/10 dark:bg-white/15 overflow-hidden">
<div
className={cn("h-full rounded-full transition-all", severity.bar)}
style={{ width: `${percent}%` }}
/>
</div>
) : null}
</button>
</TooltipTrigger>
<TooltipContent
@ -68,12 +97,14 @@ export const ContextUsageBar: FC<{
className="[&_span>svg]:hidden!"
>
<div className="grid min-w-44 gap-1.5 text-xs">
<div className="flex items-center justify-between gap-4">
<span className="text-muted-foreground">Context usage</span>
<span className={cn("font-mono tabular-nums font-medium", severity.text)}>
{percent.toFixed(1)}%
</span>
</div>
{hasKnownLimit && percent !== null ? (
<div className="flex items-center justify-between gap-4">
<span className="text-muted-foreground">Context usage</span>
<span className={cn("font-mono tabular-nums font-medium", severity.text)}>
{percent.toFixed(1)}%
</span>
</div>
) : null}
{promptTokens !== undefined && (
<div className="flex items-center justify-between gap-4">
<span className="text-muted-foreground">Prompt tokens</span>
@ -98,20 +129,32 @@ export const ContextUsageBar: FC<{
</span>
</div>
)}
{cacheWrites !== undefined && cacheWrites > 0 && (
<div className="flex items-center justify-between gap-4">
<span className="text-muted-foreground">Cache writes</span>
<span className="font-mono tabular-nums">
{formatTokenCountFull(cacheWrites)}
</span>
</div>
)}
<div className="my-0.5 border-t border-border/40" />
<div className="flex items-center justify-between gap-4">
<span className="text-muted-foreground">Total</span>
<span className="text-muted-foreground">
{hasKnownLimit ? "Total" : "Total tokens"}
</span>
<span className="font-mono tabular-nums">
{formatTokenCountFull(used)} / {formatTokenCountFull(total)}
{hasKnownLimit
? `${formatTokenCountFull(used)} / ${formatTokenCountFull(total as number)}`
: formatTokenCountFull(used)}
</span>
</div>
{percent > 85 && (
{hasKnownLimit && percent !== null && percent > 85 ? (
<div className="mt-1 max-w-64 text-[11px] leading-snug text-muted-foreground/90">
Close to the context limit. Generation will stop at 100%.
Increase <span className="font-medium">Context Length</span> in
the chat Settings panel to keep going.
</div>
)}
) : null}
</div>
</TooltipContent>
</Tooltip>

View file

@ -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:

View file

@ -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);
}

View file

@ -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({
<span>Images</span>
</button>
)}
{showWebFetchPill && (
<button
type="button"
disabled={webFetchDisabled}
onClick={() => setWebFetchToolsEnabled(!webFetchToolsEnabled)}
className="composer-pill-btn"
data-active={
webFetchToolsEnabled && !webFetchDisabled ? "true" : "false"
}
aria-label={
webFetchToolsEnabled ? "Disable URL fetch" : "Enable URL fetch"
}
>
<DownloadIcon className="size-3.5" />
<span>Fetch</span>
</button>
)}
</div>
<div className="flex items-center gap-1">
{dictationSupported && (

View file

@ -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::<providerId>::<modelId>`. 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<ChatRuntimeStore>((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<ChatRuntimeStore>((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<ChatRuntimeStore>((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<ChatRuntimeStore>((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<ChatRuntimeStore>((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) =>

View file

@ -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 {

View file

@ -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 {

View file

@ -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;
}

View file

@ -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 = "</s>"
def get_vocab(self):
return {"</s>": 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 = ["</s>"],
)
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 = ["</s>"],
)
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 = ["</s>"],
)
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

View file

@ -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