Merge remote-tracking branch 'origin/main' into pr-5717
Resolves package.json conflict: keep main's biome 1->2 simplification (`biome check` without trailing ".") and the test / test:watch scripts this branch added for vitest.
This commit is contained in:
commit
b62e4d18cd
74 changed files with 10967 additions and 1234 deletions
|
|
@ -109,7 +109,7 @@ def _apply_data_designer_image_context_patch() -> None:
|
|||
return
|
||||
|
||||
try:
|
||||
from data_designer.config.models import ImageContext
|
||||
from data_designer.config.models import ImageContext # pyright: ignore[reportMissingImports]
|
||||
except ImportError:
|
||||
return
|
||||
|
||||
|
|
@ -131,7 +131,7 @@ def _apply_data_designer_image_context_patch() -> None:
|
|||
|
||||
|
||||
def build_model_providers(recipe: dict[str, Any]):
|
||||
from data_designer.config.models import ModelProvider
|
||||
from data_designer.config.models import ModelProvider # pyright: ignore[reportMissingImports]
|
||||
|
||||
providers: list[ModelProvider] = []
|
||||
for provider in recipe.get("model_providers", []):
|
||||
|
|
@ -174,7 +174,7 @@ def _validate_recipe_runtime_support(
|
|||
def build_mcp_providers(
|
||||
recipe: dict[str, Any],
|
||||
) -> list:
|
||||
from data_designer.config.mcp import LocalStdioMCPProvider, MCPProvider
|
||||
from data_designer.config.mcp import LocalStdioMCPProvider, MCPProvider # pyright: ignore[reportMissingImports]
|
||||
|
||||
providers: list[MCPProvider | LocalStdioMCPProvider] = []
|
||||
for provider in recipe.get("mcp_providers", []):
|
||||
|
|
@ -214,16 +214,42 @@ def build_mcp_providers(
|
|||
return providers
|
||||
|
||||
|
||||
def _strip_frontend_model_config_metadata(recipe: dict[str, Any]) -> dict[str, Any]:
|
||||
model_configs = recipe.get("model_configs")
|
||||
if not isinstance(model_configs, list):
|
||||
return recipe
|
||||
|
||||
changed = False
|
||||
next_model_configs: list[Any] = []
|
||||
for model_config in model_configs:
|
||||
if isinstance(model_config, dict) and "gguf_variant" in model_config:
|
||||
next_model_config = dict(model_config)
|
||||
next_model_config.pop("gguf_variant", None)
|
||||
next_model_configs.append(next_model_config)
|
||||
changed = True
|
||||
continue
|
||||
next_model_configs.append(model_config)
|
||||
|
||||
if not changed:
|
||||
return recipe
|
||||
|
||||
return {
|
||||
**recipe,
|
||||
"model_configs": next_model_configs,
|
||||
}
|
||||
|
||||
|
||||
def build_config_builder(recipe: dict[str, Any]):
|
||||
_apply_data_designer_image_context_patch()
|
||||
from data_designer.config import DataDesignerConfigBuilder
|
||||
from data_designer.config.processors import ProcessorType
|
||||
from data_designer.config import DataDesignerConfigBuilder # pyright: ignore[reportMissingImports]
|
||||
from data_designer.config.processors import ProcessorType # pyright: ignore[reportMissingImports]
|
||||
|
||||
recipe_core = {
|
||||
key: value
|
||||
for key, value in recipe.items()
|
||||
if key not in {"model_providers", "mcp_providers"}
|
||||
}
|
||||
recipe_core = _strip_frontend_model_config_metadata(recipe_core)
|
||||
recipe_core, oxc_local_callable_specs = split_oxc_local_callable_validators(
|
||||
recipe_core
|
||||
)
|
||||
|
|
@ -256,8 +282,9 @@ def create_data_designer(
|
|||
artifact_path: str | None = None,
|
||||
):
|
||||
_apply_data_designer_image_context_patch()
|
||||
from data_designer.interface.data_designer import DataDesigner
|
||||
from data_designer.interface.data_designer import DataDesigner # pyright: ignore[reportMissingImports]
|
||||
|
||||
recipe = _strip_frontend_model_config_metadata(recipe)
|
||||
model_providers = build_model_providers(recipe)
|
||||
_validate_recipe_runtime_support(recipe, model_providers)
|
||||
|
||||
|
|
@ -265,7 +292,7 @@ def create_data_designer(
|
|||
# when the pipeline contains no LLM columns. Supply a lightweight stub
|
||||
# so sampler/expression-only recipes can run without a real provider.
|
||||
if not model_providers:
|
||||
from data_designer.config.models import ModelProvider
|
||||
from data_designer.config.models import ModelProvider # pyright: ignore[reportMissingImports]
|
||||
|
||||
model_providers = [
|
||||
ModelProvider(
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -60,7 +60,9 @@ _INTENT_SIGNAL = re.compile(
|
|||
# Handles both straight and curly apostrophes.
|
||||
# Excludes "I can", "I should", "I want to", "let's" which
|
||||
# appear frequently in direct answers / explanations.
|
||||
r"\b(i['\u2019](ll|m going to|m gonna)|i am (going to|gonna)|i will|i shall|let me|allow me)\b"
|
||||
# Negative lookahead drops negated forms ("I will not", "I'll never")
|
||||
# so a refusal doesn't trigger a re-prompt.
|
||||
r"\b(i['\u2019](ll|m going to|m gonna)|i am (going to|gonna)|i will|i shall|let me|allow me)\b(?!\s+(?:not|never)\b)"
|
||||
r"|"
|
||||
# Step/plan framing: "First ...", "Step 1:", "Here's my plan"
|
||||
r"\b(?:first\b|step \d+:?|here['\u2019]?s (?:my |the |a )?(?:plan|approach))"
|
||||
|
|
|
|||
|
|
@ -1,46 +1,29 @@
|
|||
# SPDX-License-Identifier: AGPL-3.0-only
|
||||
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
|
||||
|
||||
"""Validator for user-supplied llama-server pass-through args.
|
||||
"""Boundary validator for user-supplied llama-server pass-through args.
|
||||
|
||||
Studio runs llama-server as a managed subprocess and lets callers pass
|
||||
extra flags directly (CLI: ``unsloth run ... --top-k 20``; HTTP:
|
||||
``LoadRequest.llama_extra_args``). This module is the boundary that
|
||||
rejects only flags Studio fundamentally cannot share with the user --
|
||||
model identity, the auth key, and the network endpoint Studio's HTTP
|
||||
proxy targets. Anything else passes through.
|
||||
Reject only flags Studio manages (model identity, auth, network,
|
||||
parallel slots). Everything else (sampling, ``-c``, ``-ngl``,
|
||||
``--flash-attn``, ``--cache-type-*``, ``--spec-*``, ``--jinja``, ...)
|
||||
is appended after Studio's auto-set flags so llama.cpp's last-wins
|
||||
parser lets the user override.
|
||||
|
||||
User-supplied args are appended to ``cmd`` after Studio's auto-set
|
||||
flags, so llama.cpp's last-wins CLI parsing makes the user's value
|
||||
override the auto-set one. That covers tunable knobs the user might
|
||||
reasonably want to override -- ``-c``/``--ctx-size``,
|
||||
``-np``/``--parallel``, ``-fa``/``--flash-attn``,
|
||||
``-ngl``/``--gpu-layers``, ``-t``/``--threads``, ``-fit``/``--fit*``,
|
||||
``--cache-type-k/v``, ``--chat-template-file/-kwargs``,
|
||||
``--spec-*``, ``--jinja``/``--no-jinja``,
|
||||
``--no-context-shift``/``--context-shift``, sampling params, etc.
|
||||
|
||||
Reference: https://github.com/ggml-org/llama.cpp/blob/master/tools/server/README.md
|
||||
Ref: https://github.com/ggml-org/llama.cpp/blob/master/tools/server/README.md
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Iterable, Optional
|
||||
|
||||
# Each group is the full set of aliases (short + long) for one
|
||||
# hard-denied flag, taken from the llama-server README. If llama.cpp
|
||||
# adds a new alias for an existing denied flag, extend the relevant
|
||||
# group.
|
||||
#
|
||||
# Flags NOT in this list (e.g. -c, --parallel, --flash-attn, -ngl,
|
||||
# -t/--threads, --jinja, --no-context-shift, --fit*, --cache-type-*,
|
||||
# --chat-template-*, --spec-*) pass through and override Studio's
|
||||
# auto-set version via llama.cpp's last-wins CLI parsing.
|
||||
# Each group = every alias (short + long) of one hard-denied flag.
|
||||
# Extend the matching group when llama.cpp adds a new alias.
|
||||
_DENYLIST_GROUPS: tuple[frozenset[str], ...] = (
|
||||
# Model identity -- Studio resolves the model from LoadRequest and
|
||||
# passes -m / mmproj after downloading from HF if needed. A second
|
||||
# -m would point at a different model than the one Studio thinks
|
||||
# is loaded.
|
||||
# Parallel slots: owned by typer --parallel; a pass-through would
|
||||
# desync app.state.llama_parallel_slots from llama-server.
|
||||
frozenset({"-np", "--parallel", "--n-parallel"}),
|
||||
# Model identity: Studio resolves it from LoadRequest; a second
|
||||
# -m would load a different model than Studio thinks it loaded.
|
||||
frozenset({"-m", "--model"}),
|
||||
frozenset({"-mu", "--model-url"}),
|
||||
frozenset({"-dr", "--docker-repo"}),
|
||||
|
|
@ -51,28 +34,21 @@ _DENYLIST_GROUPS: tuple[frozenset[str], ...] = (
|
|||
frozenset({"-hft", "--hf-token"}),
|
||||
frozenset({"-mm", "--mmproj"}),
|
||||
frozenset({"-mmu", "--mmproj-url"}),
|
||||
# Networking -- Studio binds llama-server's port and reverse-proxies
|
||||
# HTTP traffic to it. Retargeting host/port/path/prefix would
|
||||
# orphan Studio's proxy and the UI would lose the server.
|
||||
# Networking: Studio binds + proxies; retargeting orphans the proxy.
|
||||
frozenset({"--host"}),
|
||||
frozenset({"--port"}),
|
||||
frozenset({"--path"}),
|
||||
frozenset({"--api-prefix"}),
|
||||
frozenset({"--reuse-port"}),
|
||||
# Auth / TLS -- Studio terminates auth at its own layer; an
|
||||
# upstream --api-key would shadow Studio's UNSLOTH_DIRECT_STREAM
|
||||
# key, and TLS on llama-server would break the local proxy hop.
|
||||
# Auth / TLS: Studio terminates auth; upstream --api-key / TLS
|
||||
# shadows Studio's key and breaks the proxy hop.
|
||||
frozenset({"--api-key"}),
|
||||
frozenset({"--api-key-file"}),
|
||||
frozenset({"--ssl-key-file"}),
|
||||
frozenset({"--ssl-cert-file"}),
|
||||
# Single-model server -- Studio runs one model per llama-server
|
||||
# process and serves its own UI. Enabling multi-model loading or
|
||||
# llama-server's built-in web UI changes the surface clients see.
|
||||
# ``--webui``/``--no-webui`` are the legacy spelling; current
|
||||
# upstream uses ``--ui``/``--no-ui`` + ``--ui-*`` companions.
|
||||
# Keep both so the denylist matches old and new llama-server
|
||||
# binaries (Studio's prebuilt vs system-llama.cpp).
|
||||
# Built-in web UI. --webui/--no-webui is the legacy spelling;
|
||||
# upstream renamed to --ui/--no-ui + --ui-*. Keep both so prebuilt
|
||||
# and system llama.cpp binaries both match.
|
||||
frozenset({"--webui", "--no-webui"}),
|
||||
frozenset({"--ui", "--no-ui"}),
|
||||
frozenset({"--ui-config"}),
|
||||
|
|
@ -82,32 +58,46 @@ _DENYLIST_GROUPS: tuple[frozenset[str], ...] = (
|
|||
frozenset({"--models-preset"}),
|
||||
frozenset({"--models-max"}),
|
||||
frozenset({"--models-autoload", "--no-models-autoload"}),
|
||||
# Server-mode flips: --embedding / --rerank restrict llama-server to
|
||||
# those endpoints, breaking Studio's /v1/chat/completions hop.
|
||||
frozenset({"--embedding", "--embeddings"}),
|
||||
frozenset({"--rerank", "--reranking"}),
|
||||
# llama-server's own built-in tools flag would silently stack on top
|
||||
# of Studio's --enable-tools / --disable-tools policy resolver.
|
||||
frozenset({"--tools"}),
|
||||
)
|
||||
|
||||
_DENYLIST: frozenset[str] = frozenset().union(*_DENYLIST_GROUPS)
|
||||
|
||||
|
||||
def _flag_name(token: str) -> Optional[str]:
|
||||
"""Return the flag name for a token, or None if it isn't a flag.
|
||||
"""Flag name for ``token``, or None if it isn't a flag.
|
||||
|
||||
Peels ``--key=value`` to the bare ``--key``. Plain numeric values
|
||||
like ``-1`` or ``-0.5`` (e.g. ``--seed -1``) are values, not flags;
|
||||
llama-server short-form flags always start with a letter.
|
||||
Peels `--key=value` to `--key`, treats `-1` / `-0.5` as values
|
||||
(llama-server shorts always start with a letter), strips
|
||||
whitespace, and normalises attached `-np8` / signed `-np-1` /
|
||||
digit-prefix-junk `-np8x` to `-np`. Mirrors the CLI's
|
||||
`_expand_attached_np_short`.
|
||||
"""
|
||||
token = token.strip()
|
||||
if not token.startswith("-") or token in {"-", "--"}:
|
||||
return None
|
||||
if len(token) >= 2 and (token[1].isdigit() or token[1] == "."):
|
||||
return None
|
||||
return token.split("=", 1)[0]
|
||||
name = token.split("=", 1)[0]
|
||||
if len(name) > 3 and name.startswith("-np"):
|
||||
suffix = name[3:]
|
||||
if suffix[0].isdigit() or (
|
||||
len(suffix) > 1 and suffix[0] in {"-", "+"} and suffix[1].isdigit()
|
||||
):
|
||||
return "-np"
|
||||
return name
|
||||
|
||||
|
||||
def validate_extra_args(args: Optional[Iterable[str]]) -> list[str]:
|
||||
"""Validate user-supplied llama-server args.
|
||||
|
||||
Returns the args as a flat list ready to extend the llama-server
|
||||
command. Raises ``ValueError`` (with the offending flag in the
|
||||
message) the moment a token resolves to a Studio-managed flag.
|
||||
"""
|
||||
"""Validate user-supplied llama-server args. Returns a flat list
|
||||
ready to extend the llama-server command; raises ``ValueError``
|
||||
naming the offending flag on the first managed token."""
|
||||
if not args:
|
||||
return []
|
||||
out: list[str] = []
|
||||
|
|
@ -124,15 +114,15 @@ def validate_extra_args(args: Optional[Iterable[str]]) -> list[str]:
|
|||
|
||||
|
||||
def is_managed_flag(flag: str) -> bool:
|
||||
"""True if ``flag`` is a Studio-managed llama-server flag."""
|
||||
return flag in _DENYLIST
|
||||
"""True if ``flag`` is Studio-managed. Normalises via ``_flag_name``
|
||||
so `-np8` / `--parallel=8` classify like the canonical tokens."""
|
||||
normalised = _flag_name(flag)
|
||||
return normalised is not None and normalised in _DENYLIST
|
||||
|
||||
|
||||
# Pass-through flags that shadow first-class ``LoadRequest`` fields
|
||||
# (max_seq_length, cache_type_kv, speculative_type,
|
||||
# chat_template_override). Stripped from inherited extras so they
|
||||
# can't last-wins-override an Apply that re-sets the same first-class
|
||||
# field.
|
||||
# Pass-through flags that shadow first-class LoadRequest fields;
|
||||
# stripped from inherited extras so they can't last-wins-override an
|
||||
# Apply that re-sets the same field.
|
||||
_CONTEXT_FLAGS: frozenset[str] = frozenset({"-c", "--ctx-size"})
|
||||
_CACHE_FLAGS: frozenset[str] = frozenset(
|
||||
{"-ctk", "--cache-type-k", "-ctv", "--cache-type-v"}
|
||||
|
|
@ -169,9 +159,8 @@ _SHADOWING_FLAGS: frozenset[str] = (
|
|||
_CONTEXT_FLAGS | _CACHE_FLAGS | _SPEC_FLAGS | _TEMPLATE_FLAGS
|
||||
)
|
||||
|
||||
# Boolean flags inside _SHADOWING_FLAGS that take no value. The
|
||||
# value-consuming heuristic in strip_shadowing_flags must skip just the
|
||||
# flag for these, never the following token.
|
||||
# Shadowing flags that take no value -- strip the flag only, never the
|
||||
# following token.
|
||||
_BOOLEAN_SHADOWING_FLAGS: frozenset[str] = frozenset(
|
||||
{"--spec-default", "--jinja", "--no-jinja"}
|
||||
)
|
||||
|
|
@ -187,14 +176,11 @@ def strip_shadowing_flags(
|
|||
) -> list[str]:
|
||||
"""Strip flags that shadow first-class Studio settings.
|
||||
|
||||
Used when the route inherits a previous load's ``llama_extra_args``
|
||||
so that an inherited ``-c 4096`` cannot override the current
|
||||
request's ``max_seq_length`` (and equivalents for cache /
|
||||
speculative / chat template). Each ``strip_*`` flag controls one
|
||||
group; the route only strips groups whose corresponding first-class
|
||||
field was actually supplied by the caller, so an inherited
|
||||
``--chat-template-file`` survives an Apply that omits both
|
||||
``llama_extra_args`` and ``chat_template_override``.
|
||||
Used when inheriting a previous load's ``llama_extra_args`` so an
|
||||
inherited `-c 4096` can't override the current `max_seq_length`
|
||||
(same for cache / spec / template). Each ``strip_*`` toggle
|
||||
controls one group; the route only strips groups whose first-class
|
||||
field the caller actually supplied.
|
||||
"""
|
||||
shadowing: set[str] = set()
|
||||
if strip_context:
|
||||
|
|
@ -216,9 +202,8 @@ def strip_shadowing_flags(
|
|||
out.append(tok)
|
||||
i += 1
|
||||
continue
|
||||
# Drop this token. Boolean shadowing flags never carry a value;
|
||||
# other shadowing flags consume the next token when it isn't a
|
||||
# flag and the value isn't already packed as ``--key=value``.
|
||||
# Drop the flag; consume the next token too unless it's
|
||||
# boolean, already inline (`-c=4096`), or another flag.
|
||||
if flag in _BOOLEAN_SHADOWING_FLAGS or "=" in tok:
|
||||
i += 1
|
||||
elif i + 1 < n and _flag_name(tokens[i + 1]) is None:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
},
|
||||
|
|
|
|||
|
|
@ -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}$")
|
||||
|
|
@ -728,10 +730,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
|
||||
|
|
@ -756,6 +756,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():
|
||||
|
|
@ -766,11 +846,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(
|
||||
|
|
@ -780,11 +867,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"}
|
||||
|
||||
|
|
@ -798,6 +885,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
|
||||
|
|
|
|||
|
|
@ -299,7 +299,11 @@ class InferenceStatusResponse(BaseModel):
|
|||
"""Current inference backend status"""
|
||||
|
||||
active_model: Optional[str] = Field(
|
||||
None, description = "Currently active model identifier"
|
||||
None, description = "Currently active model display identifier"
|
||||
)
|
||||
model_identifier: Optional[str] = Field(
|
||||
None,
|
||||
description = "Loadable identifier for the active model.",
|
||||
)
|
||||
is_vision: bool = Field(
|
||||
False, description = "Whether the active model is a vision model"
|
||||
|
|
@ -471,6 +475,40 @@ class InputDocumentContentPart(BaseModel):
|
|||
)
|
||||
|
||||
|
||||
class OpenAIReasoningContentPart(BaseModel):
|
||||
"""OpenAI Responses reasoning item paired with a tool output.
|
||||
|
||||
Reasoning models can require the previous ``reasoning`` output item
|
||||
to be replayed immediately before an ``image_generation_call`` id
|
||||
when manually managing Responses context. This part is OpenAI-only;
|
||||
routes strip it for every other provider before proxying.
|
||||
"""
|
||||
|
||||
type: Literal["reasoning"]
|
||||
id: str = Field(..., description = "OpenAI reasoning output item id.")
|
||||
summary: list[dict[str, Any]] = Field(default_factory = list)
|
||||
status: Optional[Literal["in_progress", "completed", "incomplete"]] = None
|
||||
|
||||
|
||||
class ImageGenerationCallContentPart(BaseModel):
|
||||
"""OpenAI Responses image_generation call reference.
|
||||
|
||||
OpenAI accepts prior ``image_generation_call`` items in the next
|
||||
Responses ``input`` array so follow-up prompts can edit or refine a
|
||||
generated image without resending the base64 payload. The frontend
|
||||
forwards this as a synthetic assistant content part when building
|
||||
the next OpenAI Responses request; ``external_provider`` translates
|
||||
it back to the provider-specific top-level input item.
|
||||
"""
|
||||
|
||||
type: Literal["image_generation_call"]
|
||||
id: str = Field(..., description = "OpenAI image_generation_call output item id.")
|
||||
response_id: Optional[str] = Field(
|
||||
None,
|
||||
description = "OpenAI Responses response id to use as previous_response_id for follow-up edits.",
|
||||
)
|
||||
|
||||
|
||||
class CompactionContentPart(BaseModel):
|
||||
"""Anthropic server-side compaction state, attached to an assistant
|
||||
message for round-tripping on the next turn.
|
||||
|
|
@ -504,6 +542,8 @@ ContentPart = Annotated[
|
|||
Annotated[TextContentPart, Tag("text")],
|
||||
Annotated[ImageContentPart, Tag("image_url")],
|
||||
Annotated[InputDocumentContentPart, Tag("input_document")],
|
||||
Annotated[OpenAIReasoningContentPart, Tag("reasoning")],
|
||||
Annotated[ImageGenerationCallContentPart, Tag("image_generation_call")],
|
||||
Annotated[CompactionContentPart, Tag("compaction")],
|
||||
],
|
||||
Discriminator(_content_part_discriminator),
|
||||
|
|
@ -786,6 +826,16 @@ class ChatCompletionRequest(BaseModel):
|
|||
"to auto-create."
|
||||
),
|
||||
)
|
||||
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":
|
||||
|
|
|
|||
|
|
@ -0,0 +1,5 @@
|
|||
# mlx-vlm / mlx-lm declare transformers>=5.x which conflicts with the
|
||||
# main venv's constraints.txt pin transformers==4.57.6 and forces uv to
|
||||
# backtrack unsloth. Relax to match the pin -- per-model 5.x routing
|
||||
# happens at runtime via the side-car venvs.
|
||||
transformers>=4.57.6
|
||||
|
|
@ -95,6 +95,89 @@ def _used_llm_model_aliases(recipe: dict[str, Any]) -> set[str]:
|
|||
return aliases
|
||||
|
||||
|
||||
def _used_local_model_selections(
|
||||
recipe: dict[str, Any], local_provider_names: set[str]
|
||||
) -> dict[tuple[str, str], list[str]]:
|
||||
used_aliases = _used_llm_model_aliases(recipe)
|
||||
selections: dict[tuple[str, str], list[str]] = {}
|
||||
for mc in recipe.get("model_configs", []):
|
||||
if not isinstance(mc, dict):
|
||||
continue
|
||||
alias = mc.get("alias")
|
||||
if not isinstance(alias, str) or alias not in used_aliases:
|
||||
continue
|
||||
provider = mc.get("provider")
|
||||
if not isinstance(provider, str) or provider not in local_provider_names:
|
||||
continue
|
||||
model = mc.get("model")
|
||||
target = model.strip() if isinstance(model, str) else ""
|
||||
if not target or target.lower() == "local":
|
||||
continue
|
||||
variant = mc.get("gguf_variant")
|
||||
gguf_variant = variant.strip() if isinstance(variant, str) else ""
|
||||
selections.setdefault((target, gguf_variant), []).append(alias)
|
||||
return selections
|
||||
|
||||
|
||||
def _single_used_local_model_selection(
|
||||
recipe: dict[str, Any], local_provider_names: set[str]
|
||||
) -> tuple[str, str] | None:
|
||||
selections = _used_local_model_selections(recipe, local_provider_names)
|
||||
if not selections:
|
||||
return None
|
||||
if len(selections) > 1:
|
||||
aliases = ", ".join(alias for values in selections.values() for alias in values)
|
||||
raise ValueError(
|
||||
"Recipes supports one active local model per run. "
|
||||
f"Select the same local model and GGUF variant for: {aliases}."
|
||||
)
|
||||
return next(iter(selections))
|
||||
|
||||
|
||||
def _loaded_local_model_identity() -> tuple[bool, str, str]:
|
||||
from routes.inference import get_llama_cpp_backend
|
||||
from core.inference import get_inference_backend
|
||||
|
||||
llama = get_llama_cpp_backend()
|
||||
if llama.is_loaded:
|
||||
model = str(getattr(llama, "model_identifier", "") or "").strip()
|
||||
variant = str(getattr(llama, "hf_variant", "") or "").strip()
|
||||
return True, model, variant
|
||||
|
||||
backend = get_inference_backend()
|
||||
active_model = str(getattr(backend, "active_model_name", "") or "").strip()
|
||||
if active_model:
|
||||
return True, active_model, ""
|
||||
return False, "", ""
|
||||
|
||||
|
||||
def _ensure_selected_local_model_loaded(
|
||||
recipe: dict[str, Any], local_provider_names: set[str]
|
||||
) -> None:
|
||||
model_loaded, active_model, active_variant = _loaded_local_model_identity()
|
||||
if not model_loaded:
|
||||
raise ValueError(
|
||||
"No model loaded in Chat. Load a model first, then run the recipe."
|
||||
)
|
||||
|
||||
selection = _single_used_local_model_selection(recipe, local_provider_names)
|
||||
if selection is None:
|
||||
return
|
||||
|
||||
target, gguf_variant = selection
|
||||
variant_matches = not gguf_variant or active_variant == gguf_variant
|
||||
if active_model.lower() != target.lower() or not variant_matches:
|
||||
selected = f"{target} ({gguf_variant})" if gguf_variant else target
|
||||
active = (
|
||||
f"{active_model} ({active_variant})" if active_variant else active_model
|
||||
)
|
||||
raise ValueError(
|
||||
"Selected local model is not loaded. "
|
||||
f"Selected {selected}; active {active or 'none'}. "
|
||||
"Load the selected model again, then run the recipe."
|
||||
)
|
||||
|
||||
|
||||
def _inject_local_structured_response_format(
|
||||
recipe: dict[str, Any], local_provider_names: set[str]
|
||||
) -> None:
|
||||
|
|
@ -238,24 +321,12 @@ def _inject_local_providers(recipe: dict[str, Any], request: Request) -> Optiona
|
|||
token = ""
|
||||
internal_key_id: Optional[int] = None
|
||||
if local_names & referenced_providers:
|
||||
# Verify a model is loaded.
|
||||
# NOTE: This is a point-in-time check (TOCTOU). The model could be unloaded
|
||||
# or swapped after this check but before the recipe subprocess calls /v1.
|
||||
# The inference endpoint returns a clear 400 in that case.
|
||||
#
|
||||
# Imports are deferred to avoid circular dependencies with inference modules.
|
||||
from routes.inference import get_llama_cpp_backend
|
||||
from core.inference import get_inference_backend
|
||||
|
||||
llama = get_llama_cpp_backend()
|
||||
model_loaded = llama.is_loaded
|
||||
if not model_loaded:
|
||||
backend = get_inference_backend()
|
||||
model_loaded = bool(backend.active_model_name)
|
||||
if not model_loaded:
|
||||
raise ValueError(
|
||||
"No model loaded in Chat. Load a model first, then run the recipe."
|
||||
)
|
||||
# Verify the selected local model is loaded before minting a workflow
|
||||
# key. This still remains a point-in-time singleton-backend check
|
||||
# (TOCTOU): a future generation token should bind frontend load and
|
||||
# job creation, and the inference endpoint returns a clear 400 if the
|
||||
# model is later unloaded or swapped before the subprocess calls /v1.
|
||||
_ensure_selected_local_model_loaded(recipe, local_names)
|
||||
|
||||
from auth import storage # deferred: avoids circular import
|
||||
|
||||
|
|
@ -287,12 +358,12 @@ def _inject_local_providers(recipe: dict[str, Any], request: Request) -> Optiona
|
|||
providers[i].pop("extra_body", None)
|
||||
|
||||
# Force skip_health_check on any model_config that references a local
|
||||
# provider. The local /v1/models endpoint only lists the real loaded
|
||||
# model (e.g. "unsloth/llama-3.2-1b") and not the placeholder "local"
|
||||
# that the recipe sends as the model id, so data_designer's pre-flight
|
||||
# health check would otherwise fail before the first completion call.
|
||||
# The backend route ignores the model id field in chat completions, so
|
||||
# skipping the check is safe.
|
||||
# provider. The frontend now sends the explicit selected local model id,
|
||||
# but llama-server's /v1/models response can still differ from that id
|
||||
# for local paths, cache aliases, and GGUF variant loads. The recipe run
|
||||
# has already gated on a loaded local inference backend above, so the
|
||||
# data_designer model-list health check would be redundant and can reject
|
||||
# valid local selections.
|
||||
for mc in recipe.get("model_configs", []):
|
||||
if not isinstance(mc, dict):
|
||||
continue
|
||||
|
|
@ -319,7 +390,7 @@ def _inject_local_providers(recipe: dict[str, Any], request: Request) -> Optiona
|
|||
tpl_kwargs = extra_body.get("chat_template_kwargs")
|
||||
if not isinstance(tpl_kwargs, dict):
|
||||
tpl_kwargs = {}
|
||||
tpl_kwargs.setdefault("enable_thinking", False)
|
||||
tpl_kwargs["enable_thinking"] = False
|
||||
extra_body["chat_template_kwargs"] = tpl_kwargs
|
||||
params["extra_body"] = extra_body
|
||||
|
||||
|
|
|
|||
|
|
@ -606,11 +606,16 @@ async def load_model(
|
|||
backend = get_inference_backend()
|
||||
llama_backend = get_llama_cpp_backend()
|
||||
|
||||
if request.gguf_variant:
|
||||
is_direct_gguf_request = model_identifier.lower().endswith(".gguf")
|
||||
if request.gguf_variant or is_direct_gguf_request:
|
||||
gguf_variant_matches = is_direct_gguf_request or bool(
|
||||
llama_backend.hf_variant
|
||||
and request.gguf_variant
|
||||
and llama_backend.hf_variant.lower() == request.gguf_variant.lower()
|
||||
)
|
||||
if (
|
||||
llama_backend.is_loaded
|
||||
and llama_backend.hf_variant
|
||||
and llama_backend.hf_variant.lower() == request.gguf_variant.lower()
|
||||
and gguf_variant_matches
|
||||
and llama_backend.model_identifier
|
||||
and llama_backend.model_identifier.lower() == model_identifier.lower()
|
||||
# Match runtime settings too so Apply isn't dropped (#5401).
|
||||
|
|
@ -619,7 +624,8 @@ async def load_model(
|
|||
and getattr(llama_backend, "_audio_probed", True)
|
||||
):
|
||||
logger.info(
|
||||
f"Model already loaded (GGUF): {model_log_label} variant={request.gguf_variant}, skipping reload"
|
||||
"Model already loaded (GGUF): "
|
||||
f"{model_log_label} variant={request.gguf_variant or llama_backend.hf_variant}, skipping reload"
|
||||
)
|
||||
inference_config = load_inference_config(llama_backend.model_identifier)
|
||||
|
||||
|
|
@ -1373,6 +1379,7 @@ async def get_status(
|
|||
_audio_type = getattr(llama_backend, "_audio_type", None)
|
||||
return InferenceStatusResponse(
|
||||
active_model = _display_model_id,
|
||||
model_identifier = None if _native_grant_backed else _model_id,
|
||||
is_vision = llama_backend.is_vision,
|
||||
is_gguf = True,
|
||||
gguf_variant = llama_backend.hf_variant,
|
||||
|
|
@ -1435,6 +1442,7 @@ async def get_status(
|
|||
|
||||
return InferenceStatusResponse(
|
||||
active_model = backend.active_model_name,
|
||||
model_identifier = backend.active_model_name,
|
||||
is_vision = is_vision,
|
||||
is_gguf = False,
|
||||
is_audio = is_audio,
|
||||
|
|
@ -1709,6 +1717,12 @@ def _build_external_messages(
|
|||
see ``_INPUT_DOCUMENT_PROVIDERS``). For every other provider the
|
||||
part is stripped so the unknown content type doesn't reach generic
|
||||
/chat/completions passthrough and 400 the request.
|
||||
- `reasoning`: OpenAI-only Responses reasoning item paired with a
|
||||
prior tool output. Forwarded ONLY when provider_type=="openai"
|
||||
so follow-up image edits can replay the required reasoning item.
|
||||
- `image_generation_call`: OpenAI-only Responses image reference.
|
||||
Forwarded ONLY when provider_type=="openai" so follow-up image
|
||||
edits can reference prior generated images.
|
||||
- `compaction`: Anthropic-only synthetic part (round-trips server-side
|
||||
compaction state). Forwarded ONLY when provider_type=="anthropic";
|
||||
stripped for every other provider so the unknown part doesn't
|
||||
|
|
@ -1717,6 +1731,7 @@ def _build_external_messages(
|
|||
"""
|
||||
document_provider = provider_type in _INPUT_DOCUMENT_PROVIDERS
|
||||
anthropic = provider_type == "anthropic"
|
||||
openai = provider_type == "openai"
|
||||
result = []
|
||||
for msg in messages:
|
||||
if isinstance(msg.content, str):
|
||||
|
|
@ -1737,6 +1752,30 @@ def _build_external_messages(
|
|||
"image_url": {"url": part.image_url.url},
|
||||
}
|
||||
)
|
||||
elif (
|
||||
part.type == "reasoning" and openai and msg.role == "assistant"
|
||||
):
|
||||
reasoning: dict[str, Any] = {
|
||||
"type": "reasoning",
|
||||
"id": part.id,
|
||||
"summary": part.summary,
|
||||
}
|
||||
if part.status:
|
||||
reasoning["status"] = part.status
|
||||
parts.append(reasoning)
|
||||
elif (
|
||||
part.type == "image_generation_call"
|
||||
and openai
|
||||
and msg.role == "assistant"
|
||||
):
|
||||
# ExternalProviderClient maps this onto a top-level
|
||||
# Responses input item after the current user prompt,
|
||||
# or onto `previous_response_id` when response_id is
|
||||
# available from the prior Responses turn.
|
||||
image_ref = {"type": "image_generation_call", "id": part.id}
|
||||
if getattr(part, "response_id", None):
|
||||
image_ref["response_id"] = part.response_id
|
||||
parts.append(image_ref)
|
||||
elif part.type == "input_document" and document_provider:
|
||||
# ExternalProviderClient maps this onto
|
||||
# Anthropic's `document` or OpenAI Responses'
|
||||
|
|
@ -1758,6 +1797,8 @@ def _build_external_messages(
|
|||
# provider would 400 on the unknown part, so
|
||||
# gate by provider_type.
|
||||
parts.append({"type": "compaction", "content": part.content})
|
||||
if msg.role == "assistant" and not parts:
|
||||
continue
|
||||
result.append({"role": msg.role, "content": parts})
|
||||
else:
|
||||
# Non-vision provider: strip images / documents, keep
|
||||
|
|
@ -1769,8 +1810,28 @@ def _build_external_messages(
|
|||
for p in msg.content:
|
||||
if p.type == "text":
|
||||
preserved.append({"type": "text", "text": p.text})
|
||||
elif p.type == "reasoning" and openai and msg.role == "assistant":
|
||||
reasoning: dict[str, Any] = {
|
||||
"type": "reasoning",
|
||||
"id": p.id,
|
||||
"summary": p.summary,
|
||||
}
|
||||
if p.status:
|
||||
reasoning["status"] = p.status
|
||||
preserved.append(reasoning)
|
||||
elif (
|
||||
p.type == "image_generation_call"
|
||||
and openai
|
||||
and msg.role == "assistant"
|
||||
):
|
||||
image_ref = {"type": "image_generation_call", "id": p.id}
|
||||
if getattr(p, "response_id", None):
|
||||
image_ref["response_id"] = p.response_id
|
||||
preserved.append(image_ref)
|
||||
elif p.type == "compaction" and anthropic:
|
||||
preserved.append({"type": "compaction", "content": p.content})
|
||||
if msg.role == "assistant" and not preserved:
|
||||
continue
|
||||
if len(preserved) == 1 and preserved[0]["type"] == "text":
|
||||
# Single text part collapses back to a string for
|
||||
# providers that don't accept content arrays.
|
||||
|
|
@ -1876,6 +1937,7 @@ async def _proxy_to_external_provider(
|
|||
anthropic_code_exec_container_id = payload.anthropic_code_exec_container_id,
|
||||
prompt_cache_ttl = payload.prompt_cache_ttl,
|
||||
compaction_threshold = payload.compaction_threshold,
|
||||
fast_mode = payload.fast_mode,
|
||||
stream = payload.stream,
|
||||
)
|
||||
try:
|
||||
|
|
|
|||
|
|
@ -9,6 +9,7 @@ Works independently and can be moved to any directory.
|
|||
import os
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from typing import Optional
|
||||
|
||||
# Suppress annoying C-level dependency warnings globally (e.g. SwigPyPacked)
|
||||
os.environ["PYTHONWARNINGS"] = "ignore"
|
||||
|
|
@ -512,10 +513,94 @@ _server = None
|
|||
_shutdown_event = None
|
||||
|
||||
|
||||
_DEFAULT_FRONTEND_PATH = Path(__file__).resolve().parent.parent / "frontend" / "dist"
|
||||
|
||||
|
||||
def _iter_frontend_fallback_candidates() -> "list[Path]":
|
||||
"""Yield `studio/frontend/dist` paths to try when the default is missing.
|
||||
|
||||
Covers PATH-shadowed binaries whose __file__ resolves into a
|
||||
site-packages tree that never received a vite build (e.g. plain
|
||||
`pip install unsloth` from PyPI).
|
||||
"""
|
||||
import ast
|
||||
import re
|
||||
|
||||
out: list[Path] = []
|
||||
home_str = (
|
||||
os.environ.get("UNSLOTH_STUDIO_HOME")
|
||||
or os.environ.get("STUDIO_HOME")
|
||||
or str(Path.home() / ".unsloth" / "studio")
|
||||
)
|
||||
venv_dir = Path(home_str).expanduser() / "unsloth_studio"
|
||||
# Installer venv site-packages.
|
||||
for pattern in (
|
||||
"lib/python*/site-packages/studio/frontend/dist",
|
||||
"Lib/site-packages/studio/frontend/dist",
|
||||
):
|
||||
out.extend(venv_dir.glob(pattern))
|
||||
# Editable source roots referenced from the installer venv.
|
||||
for sp_pattern in ("lib/python*/site-packages", "Lib/site-packages"):
|
||||
for sp in venv_dir.glob(sp_pattern):
|
||||
for finder in sp.glob("__editable___*_finder.py"):
|
||||
try:
|
||||
src = finder.read_text(encoding = "utf-8")
|
||||
except OSError:
|
||||
continue
|
||||
# Tolerate single- or multi-line dict literals; [^}]* still
|
||||
# rejects nested dicts, which the setuptools template never
|
||||
# emits for editable installs.
|
||||
m = re.search(
|
||||
r"^MAPPING\s*(?::[^=]*)?=\s*(\{[^}]*\})", src, re.M | re.S
|
||||
)
|
||||
if not m:
|
||||
continue
|
||||
try:
|
||||
mapping = ast.literal_eval(m.group(1))
|
||||
except (SyntaxError, ValueError):
|
||||
continue
|
||||
# Defensive: literal_eval can return a set / list / None if the
|
||||
# matched literal is not a dict (regex captures `{...}`).
|
||||
if not isinstance(mapping, dict):
|
||||
continue
|
||||
studio_pkg = mapping.get("studio")
|
||||
if studio_pkg:
|
||||
out.append(Path(studio_pkg) / "frontend" / "dist")
|
||||
return out
|
||||
|
||||
|
||||
def _resolve_frontend_path(frontend_path: Path) -> tuple[Optional[Path], list[Path]]:
|
||||
"""Pick a frontend dir that actually contains `index.html`.
|
||||
|
||||
Returns (chosen, attempted). `chosen` is None if nothing servable was
|
||||
found; `attempted` is the full ordered list for diagnostics.
|
||||
"""
|
||||
attempted: list[Path] = []
|
||||
seen: set[Path] = set()
|
||||
|
||||
def _try(p: Path) -> bool:
|
||||
try:
|
||||
key = p.resolve()
|
||||
except OSError:
|
||||
key = p
|
||||
if key in seen:
|
||||
return False
|
||||
seen.add(key)
|
||||
attempted.append(p)
|
||||
return (p / "index.html").is_file()
|
||||
|
||||
if _try(Path(frontend_path)):
|
||||
return attempted[-1], attempted
|
||||
for alt in _iter_frontend_fallback_candidates():
|
||||
if _try(alt):
|
||||
return attempted[-1], attempted
|
||||
return None, attempted
|
||||
|
||||
|
||||
def run_server(
|
||||
host: str = "127.0.0.1",
|
||||
port: int = 8888,
|
||||
frontend_path: Path = Path(__file__).resolve().parent.parent / "frontend" / "dist",
|
||||
frontend_path: Path = _DEFAULT_FRONTEND_PATH,
|
||||
silent: bool = False,
|
||||
api_only: bool = False,
|
||||
llama_parallel_slots: int = 1,
|
||||
|
|
@ -584,14 +669,48 @@ def run_server(
|
|||
print("=" * 50)
|
||||
print("")
|
||||
|
||||
# Setup frontend if path provided (skip in api-only mode)
|
||||
# Setup frontend if path provided (skip in api-only mode).
|
||||
# Falls back through alternate locations if the default lacks a built
|
||||
# dist; errors out loudly rather than silently serving 404 on `/`.
|
||||
if frontend_path and not api_only:
|
||||
if setup_frontend(app, frontend_path):
|
||||
chosen, attempted = _resolve_frontend_path(Path(frontend_path))
|
||||
if chosen is not None and setup_frontend(app, chosen):
|
||||
if not silent:
|
||||
print(f"[OK] Frontend loaded from {frontend_path}")
|
||||
# Resolve so logs always show an absolute path for support.
|
||||
try:
|
||||
display = chosen.resolve()
|
||||
except OSError:
|
||||
display = chosen
|
||||
print(f"[OK] Frontend loaded from {display}")
|
||||
else:
|
||||
if not silent:
|
||||
print(f"[WARNING] Frontend not found at {frontend_path}")
|
||||
home_str = (
|
||||
os.environ.get("UNSLOTH_STUDIO_HOME")
|
||||
or os.environ.get("STUDIO_HOME")
|
||||
or str(Path.home() / ".unsloth" / "studio")
|
||||
)
|
||||
# Windows ships the user-facing shim at $STUDIO_HOME/bin/unsloth.exe
|
||||
# (a hardlink to the venv exe); Linux/macOS use the venv binary
|
||||
# at $STUDIO_HOME/unsloth_studio/bin/unsloth.
|
||||
home = Path(home_str).expanduser()
|
||||
if sys.platform == "win32":
|
||||
installer_bin = home / "bin" / "unsloth.exe"
|
||||
else:
|
||||
installer_bin = home / "unsloth_studio" / "bin" / "unsloth"
|
||||
tried_lines = "\n".join(f" - {p}" for p in attempted) or " (none)"
|
||||
raise SystemExit(
|
||||
"[ERROR] Studio frontend build not found.\n"
|
||||
f"Tried:\n{tried_lines}\n"
|
||||
"\n"
|
||||
"Likely cause: another 'unsloth' on PATH is shadowing the "
|
||||
"installer's binary and points at a site-packages tree with "
|
||||
"no built dist.\n"
|
||||
"\n"
|
||||
"Fix one of:\n"
|
||||
f" - run the installer's binary directly: {installer_bin} studio\n"
|
||||
" - pass --frontend <path/to/studio/frontend/dist>\n"
|
||||
" - pass --api-only to skip serving the web UI\n"
|
||||
" - reinstall: curl -fsSL https://unsloth.ai/install.sh | sh"
|
||||
)
|
||||
|
||||
# Resolve once; shared by the log rewrite and the banner.
|
||||
display_host = _resolve_external_ip() if host == "0.0.0.0" else host
|
||||
|
|
@ -718,7 +837,7 @@ if __name__ == "__main__":
|
|||
parser.add_argument(
|
||||
"--frontend",
|
||||
type = str,
|
||||
default = Path(__file__).resolve().parent.parent / "frontend" / "dist",
|
||||
default = _DEFAULT_FRONTEND_PATH,
|
||||
help = "Path to frontend build",
|
||||
)
|
||||
parser.add_argument("--silent", action = "store_true", help = "Suppress output")
|
||||
|
|
@ -727,11 +846,33 @@ if __name__ == "__main__":
|
|||
action = "store_true",
|
||||
help = "API server only, no frontend (for Tauri)",
|
||||
)
|
||||
# Mirror unsloth_cli/commands/studio.py's _PARALLEL_*. Default 1
|
||||
# applies only to direct backend launches; `unsloth studio run`
|
||||
# always passes its own value (4) explicitly.
|
||||
_PARALLEL_MIN = 1
|
||||
_PARALLEL_MAX = 64
|
||||
_PARALLEL_DEFAULT_PLAIN = 1
|
||||
parser.add_argument(
|
||||
"--parallel",
|
||||
"--n-parallel",
|
||||
type = int,
|
||||
default = _PARALLEL_DEFAULT_PLAIN,
|
||||
help = (
|
||||
f"llama-server parallel decode slots ({_PARALLEL_MIN}..{_PARALLEL_MAX}). "
|
||||
f"Default {_PARALLEL_DEFAULT_PLAIN}; `unsloth studio run` uses 4."
|
||||
),
|
||||
)
|
||||
|
||||
args = parser.parse_args()
|
||||
if not _PARALLEL_MIN <= args.parallel <= _PARALLEL_MAX:
|
||||
parser.error(f"--parallel must be between {_PARALLEL_MIN} and {_PARALLEL_MAX}")
|
||||
|
||||
kwargs = dict(
|
||||
host = args.host, port = args.port, silent = args.silent, api_only = args.api_only
|
||||
host = args.host,
|
||||
port = args.port,
|
||||
silent = args.silent,
|
||||
api_only = args.api_only,
|
||||
llama_parallel_slots = args.parallel,
|
||||
)
|
||||
if args.frontend is not None:
|
||||
kwargs["frontend_path"] = Path(args.frontend)
|
||||
|
|
|
|||
353
studio/backend/tests/test_anthropic_citations.py
Normal file
353
studio/backend/tests/test_anthropic_citations.py
Normal 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
|
||||
690
studio/backend/tests/test_anthropic_citations_edge.py
Normal file
690
studio/backend/tests/test_anthropic_citations_edge.py
Normal 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("…")
|
||||
164
studio/backend/tests/test_anthropic_fast_mode_and_refusal.py
Normal file
164
studio/backend/tests/test_anthropic_fast_mode_and_refusal.py
Normal 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
|
||||
442
studio/backend/tests/test_anthropic_fast_mode_edge.py
Normal file
442
studio/backend/tests/test_anthropic_fast_mode_edge.py
Normal 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
|
||||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
248
studio/backend/tests/test_frontend_resolution.py
Normal file
248
studio/backend/tests/test_frontend_resolution.py
Normal file
|
|
@ -0,0 +1,248 @@
|
|||
# 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 frontend-dist resolver in studio/backend/run.py.
|
||||
|
||||
Loads only the relevant helpers via importlib so the test does not pull in
|
||||
uvicorn / FastAPI / unsloth's full dependency tree. Pairs with the AST-style
|
||||
test_host_defaults.py.
|
||||
"""
|
||||
|
||||
import ast
|
||||
import importlib.util
|
||||
import os
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
_RUN_PY = Path(__file__).resolve().parent.parent / "run.py"
|
||||
_REPO_STUDIO_DIR = _RUN_PY.parent.parent # studio/
|
||||
|
||||
|
||||
def _load_helpers_only():
|
||||
"""Import just the resolver helpers from run.py without executing the
|
||||
server-side imports (uvicorn, structlog, etc.)."""
|
||||
source = _RUN_PY.read_text(encoding = "utf-8")
|
||||
tree = ast.parse(source)
|
||||
keep = []
|
||||
wanted = {
|
||||
"_DEFAULT_FRONTEND_PATH",
|
||||
"_iter_frontend_fallback_candidates",
|
||||
"_resolve_frontend_path",
|
||||
}
|
||||
for node in tree.body:
|
||||
if isinstance(node, (ast.Import, ast.ImportFrom)):
|
||||
keep.append(node)
|
||||
elif isinstance(node, ast.Assign):
|
||||
names = {t.id for t in node.targets if isinstance(t, ast.Name)}
|
||||
if names & wanted:
|
||||
keep.append(node)
|
||||
elif isinstance(node, ast.FunctionDef) and node.name in wanted:
|
||||
keep.append(node)
|
||||
module = ast.Module(body = keep, type_ignores = [])
|
||||
code = compile(module, str(_RUN_PY), "exec")
|
||||
ns: dict = {"__file__": str(_RUN_PY), "__name__": "_run_helpers_test"}
|
||||
exec(code, ns)
|
||||
return ns
|
||||
|
||||
|
||||
def test_resolver_returns_none_when_nothing_exists(tmp_path, monkeypatch):
|
||||
monkeypatch.setenv("UNSLOTH_STUDIO_HOME", str(tmp_path / "no_studio"))
|
||||
monkeypatch.delenv("STUDIO_HOME", raising = False)
|
||||
helpers = _load_helpers_only()
|
||||
chosen, attempted = helpers["_resolve_frontend_path"](tmp_path / "missing")
|
||||
assert chosen is None
|
||||
assert attempted == [tmp_path / "missing"]
|
||||
|
||||
|
||||
def test_resolver_picks_first_existing_candidate(tmp_path, monkeypatch):
|
||||
dist = tmp_path / "good" / "frontend" / "dist"
|
||||
dist.mkdir(parents = True)
|
||||
(dist / "index.html").write_text("<!doctype html>", encoding = "utf-8")
|
||||
monkeypatch.setenv("UNSLOTH_STUDIO_HOME", str(tmp_path / "no_studio"))
|
||||
monkeypatch.delenv("STUDIO_HOME", raising = False)
|
||||
helpers = _load_helpers_only()
|
||||
chosen, attempted = helpers["_resolve_frontend_path"](dist)
|
||||
assert chosen == dist
|
||||
assert attempted[-1] == dist
|
||||
|
||||
|
||||
def test_resolver_falls_back_to_studio_home_site_packages(tmp_path, monkeypatch):
|
||||
studio_home = tmp_path / "studio_home"
|
||||
sp_dist = (
|
||||
studio_home
|
||||
/ "unsloth_studio"
|
||||
/ "lib"
|
||||
/ "python3.13"
|
||||
/ "site-packages"
|
||||
/ "studio"
|
||||
/ "frontend"
|
||||
/ "dist"
|
||||
)
|
||||
sp_dist.mkdir(parents = True)
|
||||
(sp_dist / "index.html").write_text("<!doctype html>", encoding = "utf-8")
|
||||
monkeypatch.setenv("UNSLOTH_STUDIO_HOME", str(studio_home))
|
||||
monkeypatch.delenv("STUDIO_HOME", raising = False)
|
||||
helpers = _load_helpers_only()
|
||||
chosen, attempted = helpers["_resolve_frontend_path"](tmp_path / "bogus")
|
||||
assert chosen is not None
|
||||
assert chosen.resolve() == sp_dist.resolve()
|
||||
assert (tmp_path / "bogus") in attempted
|
||||
|
||||
|
||||
def test_resolver_falls_back_via_editable_pth(tmp_path, monkeypatch):
|
||||
"""Simulates a `--local` install: dedicated venv with an editable .pth
|
||||
pointing at a cloned repo that owns the built dist."""
|
||||
studio_home = tmp_path / "studio_home"
|
||||
sp = studio_home / "unsloth_studio" / "lib" / "python3.13" / "site-packages"
|
||||
sp.mkdir(parents = True)
|
||||
repo_root = tmp_path / "clone"
|
||||
repo_studio = repo_root / "studio"
|
||||
repo_dist = repo_studio / "frontend" / "dist"
|
||||
repo_dist.mkdir(parents = True)
|
||||
(repo_dist / "index.html").write_text("<!doctype html>", encoding = "utf-8")
|
||||
# Minimal `__editable___pkg_finder.py` carrying a MAPPING dict that
|
||||
# setuptools' editable install generator writes.
|
||||
finder = sp / "__editable___unsloth_0_0_0_finder.py"
|
||||
finder.write_text(
|
||||
"MAPPING: dict[str, str] = "
|
||||
f"{{'studio': {str(repo_studio)!r}, 'unsloth': '/x', 'unsloth_cli': '/y'}}\n",
|
||||
encoding = "utf-8",
|
||||
)
|
||||
monkeypatch.setenv("UNSLOTH_STUDIO_HOME", str(studio_home))
|
||||
monkeypatch.delenv("STUDIO_HOME", raising = False)
|
||||
helpers = _load_helpers_only()
|
||||
chosen, attempted = helpers["_resolve_frontend_path"](tmp_path / "bogus")
|
||||
assert chosen is not None
|
||||
assert chosen.resolve() == repo_dist.resolve()
|
||||
|
||||
|
||||
def test_iter_candidates_handles_missing_studio_home(tmp_path, monkeypatch):
|
||||
monkeypatch.setenv("UNSLOTH_STUDIO_HOME", str(tmp_path / "nonexistent"))
|
||||
monkeypatch.delenv("STUDIO_HOME", raising = False)
|
||||
helpers = _load_helpers_only()
|
||||
# Glob over a non-existent dir is empty; must not raise.
|
||||
candidates = helpers["_iter_frontend_fallback_candidates"]()
|
||||
assert candidates == []
|
||||
|
||||
|
||||
def test_resolver_falls_back_to_windows_layout_site_packages(tmp_path, monkeypatch):
|
||||
"""Pins the `Lib/site-packages` (capital L) Windows venv layout
|
||||
alongside the POSIX `lib/python*/site-packages` path."""
|
||||
studio_home = tmp_path / "studio_home"
|
||||
sp_dist = (
|
||||
studio_home
|
||||
/ "unsloth_studio"
|
||||
/ "Lib"
|
||||
/ "site-packages"
|
||||
/ "studio"
|
||||
/ "frontend"
|
||||
/ "dist"
|
||||
)
|
||||
sp_dist.mkdir(parents = True)
|
||||
(sp_dist / "index.html").write_text("<!doctype html>", encoding = "utf-8")
|
||||
monkeypatch.setenv("UNSLOTH_STUDIO_HOME", str(studio_home))
|
||||
monkeypatch.delenv("STUDIO_HOME", raising = False)
|
||||
helpers = _load_helpers_only()
|
||||
chosen, _ = helpers["_resolve_frontend_path"](tmp_path / "bogus")
|
||||
assert chosen is not None
|
||||
assert chosen.resolve() == sp_dist.resolve()
|
||||
|
||||
|
||||
def test_resolver_does_not_crash_on_non_dict_mapping_literal(tmp_path, monkeypatch):
|
||||
"""A finder file whose MAPPING value is a set / list / non-dict literal
|
||||
(theoretically possible if the regex matched a brace-delimited literal
|
||||
that ast.literal_eval can parse) must not AttributeError. The resolver
|
||||
should skip that finder and keep probing."""
|
||||
studio_home = tmp_path / "studio_home"
|
||||
sp = studio_home / "unsloth_studio" / "lib" / "python3.13" / "site-packages"
|
||||
sp.mkdir(parents = True)
|
||||
# Bad finder: set literal, not a dict. ast.literal_eval parses it as set;
|
||||
# any .get() call on it would raise AttributeError.
|
||||
(sp / "__editable___bad_0_0_0_finder.py").write_text(
|
||||
"MAPPING: dict[str, str] = {'studio', 'unsloth', 'unsloth_cli'}\n",
|
||||
encoding = "utf-8",
|
||||
)
|
||||
# Good finder that should still be discovered after the bad one is skipped.
|
||||
repo_root = tmp_path / "clone"
|
||||
repo_dist = repo_root / "studio" / "frontend" / "dist"
|
||||
repo_dist.mkdir(parents = True)
|
||||
(repo_dist / "index.html").write_text("<!doctype html>", encoding = "utf-8")
|
||||
(sp / "__editable___good_0_0_0_finder.py").write_text(
|
||||
f"MAPPING: dict[str, str] = {{'studio': {str(repo_root / 'studio')!r}}}\n",
|
||||
encoding = "utf-8",
|
||||
)
|
||||
monkeypatch.setenv("UNSLOTH_STUDIO_HOME", str(studio_home))
|
||||
monkeypatch.delenv("STUDIO_HOME", raising = False)
|
||||
helpers = _load_helpers_only()
|
||||
chosen, _ = helpers["_resolve_frontend_path"](tmp_path / "bogus")
|
||||
assert chosen is not None
|
||||
assert chosen.resolve() == repo_dist.resolve()
|
||||
|
||||
|
||||
def test_resolver_handles_multiline_mapping_dict(tmp_path, monkeypatch):
|
||||
"""A future setuptools / black reformat that wraps the MAPPING dict
|
||||
across multiple lines must still parse and resolve. Locks in the
|
||||
`[^}]*` + re.DOTALL behaviour."""
|
||||
studio_home = tmp_path / "studio_home"
|
||||
sp = studio_home / "unsloth_studio" / "lib" / "python3.13" / "site-packages"
|
||||
sp.mkdir(parents = True)
|
||||
repo_root = tmp_path / "clone"
|
||||
repo_studio = repo_root / "studio"
|
||||
repo_dist = repo_studio / "frontend" / "dist"
|
||||
repo_dist.mkdir(parents = True)
|
||||
(repo_dist / "index.html").write_text("<!doctype html>", encoding = "utf-8")
|
||||
finder = sp / "__editable___unsloth_0_0_0_finder.py"
|
||||
finder.write_text(
|
||||
"MAPPING: dict[str, str] = {\n"
|
||||
f" 'studio': {str(repo_studio)!r},\n"
|
||||
" 'unsloth': '/x',\n"
|
||||
" 'unsloth_cli': '/y',\n"
|
||||
"}\n",
|
||||
encoding = "utf-8",
|
||||
)
|
||||
monkeypatch.setenv("UNSLOTH_STUDIO_HOME", str(studio_home))
|
||||
monkeypatch.delenv("STUDIO_HOME", raising = False)
|
||||
helpers = _load_helpers_only()
|
||||
chosen, _ = helpers["_resolve_frontend_path"](tmp_path / "bogus")
|
||||
assert chosen is not None
|
||||
assert chosen.resolve() == repo_dist.resolve()
|
||||
|
||||
|
||||
def test_systemexit_message_contains_actionable_fixes(tmp_path, monkeypatch):
|
||||
"""The user-facing recovery message is a contract: it must surface the
|
||||
attempted paths and every concrete fix. Pin its structure so a future
|
||||
refactor doesn't drop one."""
|
||||
import os
|
||||
import sys
|
||||
|
||||
monkeypatch.setenv("UNSLOTH_STUDIO_HOME", str(tmp_path / "no_studio"))
|
||||
monkeypatch.delenv("STUDIO_HOME", raising = False)
|
||||
helpers = _load_helpers_only()
|
||||
bogus = tmp_path / "no_such_dist"
|
||||
_, attempted = helpers["_resolve_frontend_path"](bogus)
|
||||
home = Path(os.environ["UNSLOTH_STUDIO_HOME"]).expanduser()
|
||||
if sys.platform == "win32":
|
||||
installer_bin = home / "bin" / "unsloth.exe"
|
||||
else:
|
||||
installer_bin = home / "unsloth_studio" / "bin" / "unsloth"
|
||||
tried_lines = "\n".join(f" - {p}" for p in attempted)
|
||||
message = (
|
||||
"[ERROR] Studio frontend build not found.\n"
|
||||
f"Tried:\n{tried_lines}\n"
|
||||
"\n"
|
||||
"Likely cause: another 'unsloth' on PATH is shadowing the "
|
||||
"installer's binary and points at a site-packages tree with "
|
||||
"no built dist.\n"
|
||||
"\n"
|
||||
"Fix one of:\n"
|
||||
f" - run the installer's binary directly: {installer_bin} studio\n"
|
||||
" - pass --frontend <path/to/studio/frontend/dist>\n"
|
||||
" - pass --api-only to skip serving the web UI\n"
|
||||
" - reinstall: curl -fsSL https://unsloth.ai/install.sh | sh"
|
||||
)
|
||||
assert str(bogus) in message
|
||||
assert "--frontend" in message
|
||||
assert "--api-only" in message
|
||||
assert "reinstall" in message
|
||||
assert "installer's binary directly" in message
|
||||
assert str(installer_bin) in message
|
||||
144
studio/backend/tests/test_index_bootstrap_origin.py
Normal file
144
studio/backend/tests/test_index_bootstrap_origin.py
Normal 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
|
||||
196
studio/backend/tests/test_index_bootstrap_origin_extra.py
Normal file
196
studio/backend/tests/test_index_bootstrap_origin_extra.py
Normal 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
|
||||
|
|
@ -3,21 +3,35 @@
|
|||
|
||||
"""Unit tests for the llama-server pass-through args validator.
|
||||
|
||||
The validator is the security boundary between user-supplied CLI / HTTP
|
||||
input and the llama-server subprocess command. These tests pin the
|
||||
denylist behavior so the boundary doesn't quietly regress when new
|
||||
managed flags are added.
|
||||
The validator is the boundary between user CLI/HTTP input and the
|
||||
llama-server subprocess. These tests pin denylist behaviour so it
|
||||
doesn't quietly regress when new managed flags are added.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import importlib.util
|
||||
import re
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
from core.inference.llama_server_args import (
|
||||
is_managed_flag,
|
||||
strip_shadowing_flags,
|
||||
validate_extra_args,
|
||||
# Load llama_server_args.py directly so this test doesn't drag in the
|
||||
# full backend chain (fastapi / structlog / loggers / utils.hardware)
|
||||
# via core/inference/__init__.py. The validator is intentionally
|
||||
# dependency-free and unit-tests should reflect that.
|
||||
_LSA_PATH = (
|
||||
Path(__file__).resolve().parent.parent
|
||||
/ "core"
|
||||
/ "inference"
|
||||
/ "llama_server_args.py"
|
||||
)
|
||||
_spec = importlib.util.spec_from_file_location("_lsa_test_only", _LSA_PATH)
|
||||
_lsa = importlib.util.module_from_spec(_spec)
|
||||
_spec.loader.exec_module(_lsa)
|
||||
is_managed_flag = _lsa.is_managed_flag
|
||||
strip_shadowing_flags = _lsa.strip_shadowing_flags
|
||||
validate_extra_args = _lsa.validate_extra_args
|
||||
|
||||
|
||||
# ── Pass-through (allowed) ───────────────────────────────────────────
|
||||
|
|
@ -60,13 +74,12 @@ from core.inference.llama_server_args import (
|
|||
# Reasoning controls
|
||||
["--reasoning-format", "deepseek"],
|
||||
["-rea", "auto"],
|
||||
# Soft-managed flags the user may want to override on the CLI;
|
||||
# llama.cpp's last-wins parsing means these win over Studio's
|
||||
# auto-set version.
|
||||
# Soft-managed: user-supplied flags last-wins-override Studio's
|
||||
# auto-set version. --parallel / -np / --n-parallel are NOT
|
||||
# here -- they're hard-denied (KV-cache + slot count would
|
||||
# desync). Use `unsloth studio run --parallel N` instead.
|
||||
["-c", "131072"],
|
||||
["--ctx-size", "8192"],
|
||||
["--parallel", "1"],
|
||||
["-np", "8"],
|
||||
["--flash-attn", "off"],
|
||||
["-fa", "on"],
|
||||
["--no-context-shift"],
|
||||
|
|
@ -99,8 +112,7 @@ def test_value_with_equals_form_passes_through():
|
|||
|
||||
|
||||
def test_non_flag_token_passes_through():
|
||||
# A bare positional value (not preceded by a flag) is preserved
|
||||
# verbatim. llama-server may reject it, but that's not our job.
|
||||
# Bare positionals are passed through; llama-server can reject them.
|
||||
assert validate_extra_args(["foo"]) == ["foo"]
|
||||
|
||||
|
||||
|
|
@ -110,18 +122,33 @@ def test_non_flag_token_passes_through():
|
|||
@pytest.mark.parametrize(
|
||||
"denied",
|
||||
[
|
||||
# Model identity
|
||||
# Parallel slots -- owned by the typer --parallel flag.
|
||||
"-np",
|
||||
"--parallel",
|
||||
"--n-parallel",
|
||||
# Model identity (every alias; bumping llama.cpp must keep
|
||||
# every form rejected, not just the long).
|
||||
"-m",
|
||||
"--model",
|
||||
"-mu",
|
||||
"--model-url",
|
||||
"-dr",
|
||||
"--docker-repo",
|
||||
"-hf",
|
||||
"-hfr",
|
||||
"--hf-repo",
|
||||
"-hff",
|
||||
"--hf-file",
|
||||
"-hfv",
|
||||
"-hfrv",
|
||||
"--hf-repo-v",
|
||||
"-hffv",
|
||||
"--hf-file-v",
|
||||
"-hft",
|
||||
"--hf-token",
|
||||
"-mm",
|
||||
"--mmproj",
|
||||
"-mmu",
|
||||
"--mmproj-url",
|
||||
# Networking (Studio binds + proxies)
|
||||
"--host",
|
||||
|
|
@ -134,11 +161,28 @@ def test_non_flag_token_passes_through():
|
|||
"--api-key-file",
|
||||
"--ssl-key-file",
|
||||
"--ssl-cert-file",
|
||||
# Single-model server
|
||||
# Single-model server (legacy --webui + current --ui group)
|
||||
"--webui",
|
||||
"--no-webui",
|
||||
"--ui",
|
||||
"--no-ui",
|
||||
"--ui-config",
|
||||
"--ui-config-file",
|
||||
"--ui-mcp-proxy",
|
||||
"--no-ui-mcp-proxy",
|
||||
"--models-dir",
|
||||
"--models-preset",
|
||||
"--models-max",
|
||||
"--models-autoload",
|
||||
"--no-models-autoload",
|
||||
# Server-mode flips: --embedding / --rerank would restrict
|
||||
# llama-server to those endpoints and break Studio's chat hop.
|
||||
"--embedding",
|
||||
"--embeddings",
|
||||
"--rerank",
|
||||
"--reranking",
|
||||
# llama-server's own --tools clashes with Studio's tool policy.
|
||||
"--tools",
|
||||
],
|
||||
)
|
||||
def test_denylist_rejects_all_aliases(denied):
|
||||
|
|
@ -146,14 +190,65 @@ def test_denylist_rejects_all_aliases(denied):
|
|||
validate_extra_args([denied, "value"])
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"args,offending",
|
||||
[
|
||||
# Pass-through --parallel would last-wins-override the real
|
||||
# slot count while Studio's KV-cache fit + llama_parallel_slots
|
||||
# stay at the typer value -- plan vs. process disagree.
|
||||
(["--parallel", "8"], "--parallel"),
|
||||
(["--parallel=8"], "--parallel"),
|
||||
(["--n-parallel", "16"], "--n-parallel"),
|
||||
(["--n-parallel=16"], "--n-parallel"),
|
||||
(["-np", "32"], "-np"),
|
||||
# Attached short form: Click clusters it CLI-side; HTTP /load
|
||||
# with `["-np8"]` must still resolve to managed.
|
||||
(["-np8"], "-np"),
|
||||
(["-np64"], "-np"),
|
||||
# Out-of-range values that would bypass the typer 1..64 guard.
|
||||
(["--parallel", "999"], "--parallel"),
|
||||
(["-np", "0"], "-np"),
|
||||
(["-np999"], "-np"),
|
||||
# Signed attached forms; `-np-1` must not slip past.
|
||||
(["-np-1"], "-np"),
|
||||
(["-np+1"], "-np"),
|
||||
],
|
||||
)
|
||||
def test_parallel_flags_are_managed(args, offending):
|
||||
with pytest.raises(ValueError, match = re.escape(offending)):
|
||||
validate_extra_args(args)
|
||||
|
||||
|
||||
def test_denylist_rejects_equals_form():
|
||||
with pytest.raises(ValueError, match = "--port"):
|
||||
validate_extra_args(["--port=9000"])
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"padded",
|
||||
[" --parallel", "--parallel ", "\t--parallel", " -np", "-np \n", "-np\t"],
|
||||
)
|
||||
def test_denylist_rejects_whitespace_padded_forms(padded):
|
||||
# `_flag_name` trims whitespace before lookup; otherwise a trailing
|
||||
# space could slip a managed flag past the boundary.
|
||||
with pytest.raises(ValueError, match = "parallel|np"):
|
||||
validate_extra_args([padded, "8"])
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"attached",
|
||||
["-np8x", "-np-1foo", "-np+1bar", "-np9zzz"],
|
||||
)
|
||||
def test_denylist_rejects_np_with_digit_prefix_and_junk(attached):
|
||||
# Backend `_flag_name` must classify the same forms the CLI
|
||||
# rewriter expands, else HTTP /load could smuggle `-np8x` through.
|
||||
with pytest.raises(ValueError, match = "np"):
|
||||
validate_extra_args([attached])
|
||||
|
||||
|
||||
def test_denylist_rejects_short_form_when_long_is_denied():
|
||||
# -m is the short form of the hard-denied --model; rejecting only
|
||||
# the long form would leave a trivial bypass.
|
||||
# `-m` is the short form of --model; rejecting only the long
|
||||
# form would leave a trivial bypass.
|
||||
with pytest.raises(ValueError, match = "-m"):
|
||||
validate_extra_args(["-m", "/some/other/path.gguf"])
|
||||
|
||||
|
|
@ -165,9 +260,7 @@ def test_denylist_message_names_offending_flag():
|
|||
|
||||
|
||||
def test_first_denied_flag_short_circuits():
|
||||
# Validation stops at the first denied flag; later denied flags
|
||||
# in the same call don't matter for behaviour, but the message
|
||||
# should name the first one we hit.
|
||||
# Validation stops at the first denied flag; the message names it.
|
||||
with pytest.raises(ValueError, match = "--port"):
|
||||
validate_extra_args(["--port", "1", "--host", "x"])
|
||||
|
||||
|
|
@ -177,8 +270,7 @@ def test_first_denied_flag_short_circuits():
|
|||
|
||||
@pytest.mark.parametrize("value", ["-1", "-0.5", "-42", "-.5"])
|
||||
def test_negative_number_value_is_not_flag(value):
|
||||
# ``--seed -1`` is a value, not a flag. Validator must not try
|
||||
# to look up "-1" in the denylist.
|
||||
# `--seed -1`: the -1 is a value, not a flag.
|
||||
assert validate_extra_args(["--seed", value]) == ["--seed", value]
|
||||
|
||||
|
||||
|
|
@ -190,6 +282,15 @@ def test_is_managed_flag_true_for_denied():
|
|||
assert is_managed_flag("--api-key") is True
|
||||
assert is_managed_flag("-m") is True
|
||||
assert is_managed_flag("--model") is True
|
||||
# Parallel slots owned by the typer --parallel flag.
|
||||
assert is_managed_flag("--parallel") is True
|
||||
assert is_managed_flag("--n-parallel") is True
|
||||
assert is_managed_flag("-np") is True
|
||||
# Normalised forms must classify like the canonical token so
|
||||
# is_managed_flag filtering stays in sync with validate_extra_args.
|
||||
assert is_managed_flag("-np8") is True
|
||||
assert is_managed_flag("--parallel=8") is True
|
||||
assert is_managed_flag("--port=9000") is True
|
||||
|
||||
|
||||
def test_is_managed_flag_false_for_pass_through():
|
||||
|
|
@ -199,7 +300,6 @@ def test_is_managed_flag_false_for_pass_through():
|
|||
# Soft-managed flags pass through (last-wins override)
|
||||
assert is_managed_flag("-c") is False
|
||||
assert is_managed_flag("--ctx-size") is False
|
||||
assert is_managed_flag("--parallel") is False
|
||||
assert is_managed_flag("--flash-attn") is False
|
||||
assert is_managed_flag("-ngl") is False
|
||||
assert is_managed_flag("--threads") is False
|
||||
|
|
@ -231,8 +331,8 @@ def test_strip_shadowing_flags_keeps_context_when_not_requested():
|
|||
|
||||
|
||||
def test_strip_shadowing_flags_keeps_chat_template_when_template_disabled():
|
||||
# Caller did not supply chat_template_override; the inherited
|
||||
# --chat-template-file must survive the strip.
|
||||
# No chat_template_override supplied; inherited
|
||||
# --chat-template-file must survive.
|
||||
out = strip_shadowing_flags(
|
||||
["--chat-template-file", "/tmp/custom.jinja", "--top-k", "20"],
|
||||
strip_context = True,
|
||||
|
|
@ -282,7 +382,7 @@ def test_strip_shadowing_flags_keeps_spec_when_spec_disabled():
|
|||
|
||||
|
||||
def test_strip_shadowing_flags_drops_mtp_flags_when_requested():
|
||||
# MTP / draft-mtp flags must be stripped when speculative_type is re-applied.
|
||||
# MTP / draft-mtp flags must drop when speculative_type re-applies.
|
||||
out = strip_shadowing_flags(
|
||||
[
|
||||
"--spec-type",
|
||||
|
|
@ -311,8 +411,7 @@ def test_is_managed_flag_false_for_mtp_pass_through():
|
|||
|
||||
|
||||
def test_strip_shadowing_flags_boolean_does_not_consume_next_token():
|
||||
# --spec-default is a boolean shadowing flag; the value-skipping
|
||||
# heuristic must skip just the flag, not the following positional.
|
||||
# `--spec-default` is boolean; drop just the flag, keep the next token.
|
||||
out = strip_shadowing_flags(["--spec-default", "ngram-mod"], strip_spec = True)
|
||||
assert out == ["ngram-mod"]
|
||||
|
||||
|
|
@ -343,8 +442,8 @@ def test_strip_shadowing_flags_handles_empty_input():
|
|||
|
||||
|
||||
def test_strip_shadowing_flags_defaults_strip_everything():
|
||||
# The route's already-loaded comparator calls strip_shadowing_flags
|
||||
# with no kwargs to detect ANY shadowing flag in stored extras.
|
||||
# The route's already-loaded comparator calls with no kwargs to
|
||||
# detect ANY shadowing flag in stored extras.
|
||||
out = strip_shadowing_flags(
|
||||
["-c", "4096", "--cache-type-k", "q8_0", "--spec-default", "--jinja"]
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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},
|
||||
}
|
||||
|
||||
|
||||
|
|
|
|||
251
studio/backend/tests/test_openai_citation_markers.py
Normal file
251
studio/backend/tests/test_openai_citation_markers.py
Normal 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
|
||||
413
studio/backend/tests/test_openai_citation_markers_edge.py
Normal file
413
studio/backend/tests/test_openai_citation_markers_edge.py
Normal 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
|
||||
|
|
@ -210,6 +210,7 @@ def test_image_generation_done_emits_tool_event_chunks(monkeypatch):
|
|||
assert starts[0]["arguments"] == {
|
||||
"kind": "image",
|
||||
"prompt": "A photorealistic cat sitting",
|
||||
"openai_image_generation_call_id": "img_abc",
|
||||
}
|
||||
assert ends[0]["image_b64"] == "AAAA"
|
||||
assert ends[0]["image_mime"] == "image/png"
|
||||
|
|
|
|||
372
studio/backend/tests/test_openai_tool_result_fallbacks.py
Normal file
372
studio/backend/tests/test_openai_tool_result_fallbacks.py
Normal file
|
|
@ -0,0 +1,372 @@
|
|||
# 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 tests for OpenAI Responses tool-result rendering.
|
||||
|
||||
Covers two bug classes: empty web_search cards (per-card result seeded
|
||||
with "Searching: <query>") and orphan shell_call cards (bundled-output
|
||||
fallback + final flush at response.completed / response.incomplete).
|
||||
"""
|
||||
|
||||
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)
|
||||
|
||||
|
||||
async def _collect(agen):
|
||||
out = []
|
||||
async for line in agen:
|
||||
out.append(line)
|
||||
return out
|
||||
|
||||
|
||||
def _mock_http_client(monkeypatch, handler):
|
||||
transport = httpx.MockTransport(handler)
|
||||
monkeypatch.setattr(ep_mod, "_http_client", httpx.AsyncClient(transport = transport))
|
||||
|
||||
|
||||
def _make_client(base_url: str = "https://api.openai.com/v1") -> ExternalProviderClient:
|
||||
return ExternalProviderClient(
|
||||
provider_type = "openai",
|
||||
base_url = base_url,
|
||||
api_key = "sk-test",
|
||||
)
|
||||
|
||||
|
||||
def _openai_sse(events: list[dict]) -> bytes:
|
||||
chunks: list[str] = []
|
||||
for event in events:
|
||||
chunks.append(f"event: {event['type']}")
|
||||
chunks.append(f"data: {json.dumps(event)}")
|
||||
chunks.append("")
|
||||
return ("\n".join(chunks) + "\n").encode("utf-8")
|
||||
|
||||
|
||||
def _tool_events(lines: list[str]) -> list[dict]:
|
||||
out: list[dict] = []
|
||||
for line in lines:
|
||||
if not line.startswith("data:"):
|
||||
continue
|
||||
raw = line[len("data:") :].strip()
|
||||
if not raw or raw == "[DONE]":
|
||||
continue
|
||||
try:
|
||||
parsed = json.loads(raw)
|
||||
except json.JSONDecodeError:
|
||||
continue
|
||||
if isinstance(parsed, dict) and "_toolEvent" in parsed:
|
||||
out.append(parsed["_toolEvent"])
|
||||
return out
|
||||
|
||||
|
||||
def _drive_stream(sse_events, enabled_tools, monkeypatch):
|
||||
def handler(request):
|
||||
return httpx.Response(
|
||||
200,
|
||||
content = _openai_sse(sse_events),
|
||||
headers = {"content-type": "text/event-stream"},
|
||||
)
|
||||
|
||||
_mock_http_client(monkeypatch, handler)
|
||||
|
||||
async def run():
|
||||
client = _make_client()
|
||||
return await _collect(
|
||||
client._stream_openai_responses(
|
||||
messages = [{"role": "user", "content": "x"}],
|
||||
model = "gpt-5.5",
|
||||
temperature = 0.7,
|
||||
top_p = 0.95,
|
||||
max_tokens = 4096,
|
||||
enable_thinking = None,
|
||||
reasoning_effort = None,
|
||||
enabled_tools = enabled_tools,
|
||||
)
|
||||
)
|
||||
|
||||
return _drive(run())
|
||||
|
||||
|
||||
# ── web_search per-card result ─────────────────────────────────────────
|
||||
|
||||
|
||||
def test_web_search_each_call_carries_its_own_query_as_result(monkeypatch):
|
||||
"""Each card carries its own `Searching: <query>` text; no empties."""
|
||||
sse_events = [
|
||||
{
|
||||
"type": "response.output_item.done",
|
||||
"item": {
|
||||
"type": "web_search_call",
|
||||
"id": "ws_1",
|
||||
"action": {"query": "popular animals 2026"},
|
||||
},
|
||||
},
|
||||
{
|
||||
"type": "response.output_item.done",
|
||||
"item": {
|
||||
"type": "web_search_call",
|
||||
"id": "ws_2",
|
||||
"action": {"query": "most loved animals poll"},
|
||||
},
|
||||
},
|
||||
{
|
||||
"type": "response.output_item.done",
|
||||
"item": {
|
||||
"type": "web_search_call",
|
||||
"id": "ws_3",
|
||||
"action": {"query": "tiger ranking"},
|
||||
},
|
||||
},
|
||||
{"type": "response.completed", "response": {}},
|
||||
]
|
||||
lines = _drive_stream(sse_events, ["web_search"], monkeypatch)
|
||||
events = _tool_events(lines)
|
||||
ends = [e for e in events if e["type"] == "tool_end"]
|
||||
by_id = {e["tool_call_id"]: e for e in ends}
|
||||
assert by_id["ws_1"]["result"] == "Searching: popular animals 2026"
|
||||
assert by_id["ws_2"]["result"] == "Searching: most loved animals poll"
|
||||
assert by_id["ws_3"]["result"] == "Searching: tiger ranking"
|
||||
|
||||
|
||||
def test_web_search_last_call_overwritten_with_citations(monkeypatch):
|
||||
"""Last call still gets the aggregated citation list; earlier calls
|
||||
keep their per-call `Searching:` text."""
|
||||
sse_events = [
|
||||
{
|
||||
"type": "response.output_item.done",
|
||||
"item": {
|
||||
"type": "web_search_call",
|
||||
"id": "ws_1",
|
||||
"action": {"query": "first query"},
|
||||
},
|
||||
},
|
||||
{
|
||||
"type": "response.output_item.done",
|
||||
"item": {
|
||||
"type": "web_search_call",
|
||||
"id": "ws_2",
|
||||
"action": {"query": "second query"},
|
||||
},
|
||||
},
|
||||
{
|
||||
"type": "response.output_text.annotation.added",
|
||||
"annotation": {
|
||||
"type": "url_citation",
|
||||
"url": "https://example.com/a",
|
||||
"title": "Example A",
|
||||
},
|
||||
},
|
||||
{"type": "response.completed", "response": {}},
|
||||
]
|
||||
lines = _drive_stream(sse_events, ["web_search"], monkeypatch)
|
||||
events = _tool_events(lines)
|
||||
ends = [e for e in events if e["type"] == "tool_end"]
|
||||
by_id: dict = {}
|
||||
# Keep the LAST tool_end per id (the citation overwrite for ws_2).
|
||||
for e in ends:
|
||||
by_id[e["tool_call_id"]] = e
|
||||
# First call keeps its own query.
|
||||
assert by_id["ws_1"]["result"] == "Searching: first query"
|
||||
# Last call gets overwritten with the citation block.
|
||||
assert "Title: Example A" in by_id["ws_2"]["result"]
|
||||
assert "URL: https://example.com/a" in by_id["ws_2"]["result"]
|
||||
|
||||
|
||||
def test_web_search_empty_query_falls_back_to_empty_result(monkeypatch):
|
||||
"""No query -> empty result (no `Searching:` placeholder)."""
|
||||
sse_events = [
|
||||
{
|
||||
"type": "response.output_item.done",
|
||||
"item": {
|
||||
"type": "web_search_call",
|
||||
"id": "ws_only",
|
||||
"action": {},
|
||||
},
|
||||
},
|
||||
{"type": "response.completed", "response": {}},
|
||||
]
|
||||
lines = _drive_stream(sse_events, ["web_search"], monkeypatch)
|
||||
events = _tool_events(lines)
|
||||
ends = [e for e in events if e["type"] == "tool_end"]
|
||||
assert len(ends) == 1
|
||||
assert ends[0]["result"] == ""
|
||||
|
||||
|
||||
# ── shell_call output fallbacks ────────────────────────────────────────
|
||||
|
||||
|
||||
def test_shell_call_emits_tool_end_when_output_bundled_on_done(monkeypatch):
|
||||
"""Output bundled on the shell_call done event emits tool_end."""
|
||||
sse_events = [
|
||||
{
|
||||
"type": "response.output_item.done",
|
||||
"item": {
|
||||
"type": "shell_call",
|
||||
"id": "scall_bundled",
|
||||
"action": {"commands": ["echo hi"]},
|
||||
"output": [
|
||||
{
|
||||
"stdout": "hi\n",
|
||||
"stderr": "",
|
||||
"outcome": {"type": "exit", "exit_code": 0},
|
||||
}
|
||||
],
|
||||
},
|
||||
},
|
||||
{"type": "response.completed", "response": {}},
|
||||
]
|
||||
lines = _drive_stream(sse_events, ["code_execution"], monkeypatch)
|
||||
events = _tool_events(lines)
|
||||
starts = [e for e in events if e["type"] == "tool_start"]
|
||||
ends = [e for e in events if e["type"] == "tool_end"]
|
||||
assert len(starts) == 1
|
||||
assert starts[0]["tool_call_id"] == "scall_bundled"
|
||||
assert len(ends) == 1
|
||||
assert ends[0]["tool_call_id"] == "scall_bundled"
|
||||
assert "hi" in ends[0]["result"]
|
||||
|
||||
|
||||
def test_shell_call_bundled_then_separate_output_does_not_double_emit(monkeypatch):
|
||||
"""Separate shell_call_output after bundled-output is a no-op."""
|
||||
sse_events = [
|
||||
{
|
||||
"type": "response.output_item.done",
|
||||
"item": {
|
||||
"type": "shell_call",
|
||||
"id": "scall_both",
|
||||
"action": {"commands": ["echo bundle"]},
|
||||
"output": [
|
||||
{
|
||||
"stdout": "bundle\n",
|
||||
"stderr": "",
|
||||
"outcome": {"type": "exit", "exit_code": 0},
|
||||
}
|
||||
],
|
||||
},
|
||||
},
|
||||
{
|
||||
"type": "response.output_item.done",
|
||||
"item": {
|
||||
"type": "shell_call_output",
|
||||
"id": "scout_both",
|
||||
"call_id": "scall_both",
|
||||
"output": [
|
||||
{
|
||||
"stdout": "should not double-emit\n",
|
||||
"stderr": "",
|
||||
"outcome": {"type": "exit", "exit_code": 0},
|
||||
}
|
||||
],
|
||||
},
|
||||
},
|
||||
{"type": "response.completed", "response": {}},
|
||||
]
|
||||
lines = _drive_stream(sse_events, ["code_execution"], monkeypatch)
|
||||
events = _tool_events(lines)
|
||||
ends = [e for e in events if e["type"] == "tool_end"]
|
||||
assert len(ends) == 1
|
||||
assert ends[0]["tool_call_id"] == "scall_both"
|
||||
assert "bundle" in ends[0]["result"]
|
||||
assert "should not double-emit" not in ends[0]["result"]
|
||||
|
||||
|
||||
def test_shell_call_final_flush_on_completed_when_no_output_event(monkeypatch):
|
||||
"""Orphan shell_call finalises via the response.completed flush."""
|
||||
sse_events = [
|
||||
{
|
||||
"type": "response.output_item.added",
|
||||
"item": {
|
||||
"type": "shell_call",
|
||||
"id": "scall_orphan",
|
||||
"action": {"commands": ["true"]},
|
||||
},
|
||||
},
|
||||
{
|
||||
"type": "response.output_item.done",
|
||||
"item": {
|
||||
"type": "shell_call",
|
||||
"id": "scall_orphan",
|
||||
"action": {"commands": ["true"]},
|
||||
"status": "completed",
|
||||
},
|
||||
},
|
||||
{"type": "response.completed", "response": {}},
|
||||
]
|
||||
lines = _drive_stream(sse_events, ["code_execution"], monkeypatch)
|
||||
events = _tool_events(lines)
|
||||
ends = [e for e in events if e["type"] == "tool_end"]
|
||||
assert any(e["tool_call_id"] == "scall_orphan" for e in ends)
|
||||
|
||||
|
||||
def test_shell_call_flushed_on_response_incomplete_truncation(monkeypatch):
|
||||
"""Truncated streams (response.incomplete) also flush orphan calls."""
|
||||
sse_events = [
|
||||
{
|
||||
"type": "response.output_item.added",
|
||||
"item": {
|
||||
"type": "shell_call",
|
||||
"id": "scall_truncated",
|
||||
"action": {"commands": ["long_running"]},
|
||||
},
|
||||
},
|
||||
{
|
||||
"type": "response.output_item.done",
|
||||
"item": {
|
||||
"type": "shell_call",
|
||||
"id": "scall_truncated",
|
||||
"action": {"commands": ["long_running"]},
|
||||
"status": "in_progress",
|
||||
},
|
||||
},
|
||||
{
|
||||
"type": "response.incomplete",
|
||||
"response": {
|
||||
"incomplete_details": {"reason": "max_output_tokens"},
|
||||
},
|
||||
},
|
||||
]
|
||||
lines = _drive_stream(sse_events, ["code_execution"], monkeypatch)
|
||||
events = _tool_events(lines)
|
||||
ends = [e for e in events if e["type"] == "tool_end"]
|
||||
assert any(e["tool_call_id"] == "scall_truncated" for e in ends)
|
||||
|
||||
|
||||
def test_shell_call_incomplete_does_not_double_emit(monkeypatch):
|
||||
"""response.incomplete is idempotent against already-finalised calls."""
|
||||
sse_events = [
|
||||
{
|
||||
"type": "response.output_item.done",
|
||||
"item": {
|
||||
"type": "shell_call",
|
||||
"id": "scall_done",
|
||||
"action": {"commands": ["echo done"]},
|
||||
"output": [
|
||||
{
|
||||
"stdout": "done\n",
|
||||
"stderr": "",
|
||||
"outcome": {"type": "exit", "exit_code": 0},
|
||||
}
|
||||
],
|
||||
},
|
||||
},
|
||||
{
|
||||
"type": "response.incomplete",
|
||||
"response": {
|
||||
"incomplete_details": {"reason": "max_output_tokens"},
|
||||
},
|
||||
},
|
||||
]
|
||||
lines = _drive_stream(sse_events, ["code_execution"], monkeypatch)
|
||||
events = _tool_events(lines)
|
||||
ends = [e for e in events if e["type"] == "tool_end"]
|
||||
assert len(ends) == 1
|
||||
assert ends[0]["tool_call_id"] == "scall_done"
|
||||
assert "done" in ends[0]["result"]
|
||||
|
|
@ -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
|
||||
|
|
|
|||
475
studio/backend/tests/test_pricing_edge.py
Normal file
475
studio/backend/tests/test_pricing_edge.py
Normal 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
|
||||
|
|
@ -144,7 +144,10 @@ def detect_hardware() -> DeviceType:
|
|||
if is_apple_silicon() and _has_mlx():
|
||||
DEVICE = DeviceType.MLX
|
||||
CHAT_ONLY = False
|
||||
chip = platform.processor() or platform.machine()
|
||||
# platform.processor() runs `uname -p` which returns "i386" on most
|
||||
# universal2 / Rosetta-shaped Python builds even on native arm64.
|
||||
# platform.machine() is "arm64" once is_apple_silicon() has gated us.
|
||||
chip = platform.machine() or "arm64"
|
||||
print(f"Hardware detected: MLX — Apple Silicon ({chip})")
|
||||
return DEVICE
|
||||
|
||||
|
|
@ -279,13 +282,11 @@ def get_gpu_memory_info() -> Dict[str, Any]:
|
|||
|
||||
try:
|
||||
info = mx.device_info()
|
||||
gpu_name = (
|
||||
info.get("device_name")
|
||||
or platform.processor()
|
||||
or platform.machine()
|
||||
)
|
||||
# See detect_hardware(): platform.processor() can return "i386"
|
||||
# on native arm64 Python builds, so prefer machine() as fallback.
|
||||
gpu_name = info.get("device_name") or platform.machine() or "arm64"
|
||||
except Exception:
|
||||
gpu_name = platform.processor() or platform.machine()
|
||||
gpu_name = platform.machine() or "arm64"
|
||||
|
||||
return {
|
||||
"available": True,
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue