Compare commits

..
Author SHA1 Message Date
chengyongru 720f14661f refactor(providers): declare Responses capabilities 2026-08-01 12:13:01 +08:00
11 changed files with 552 additions and 214 deletions
-7
View File
@@ -49,13 +49,6 @@ Use `/model` to inspect the current runtime model:
The response shows the current session's model and preset, plus the available preset names. Named presets come from the top-level `modelPresets` config and are the recommended way to configure model choices. `default` is always available and represents the model settings from direct `agents.defaults.*` fields. The response shows the current session's model and preset, plus the available preset names. Named presets come from the top-level `modelPresets` config and are the recommended way to configure model choices. `default` is always available and represents the model settings from direct `agents.defaults.*` fields.
`/model <preset>` expects one of those preset names, not a provider model ID or
the preset's display label. For example, if `modelPresets.local` uses the Ollama
model `llama3.2`, run `/model local`, not `/model llama3.2`. If a model is currently
configured only as an inline fallback, save it as a named preset before selecting
it manually. Fallback order controls automatic failover; it is not a list of raw
model IDs accepted by `/model`.
To switch presets for future turns: To switch presets for future turns:
```text ```text
-13
View File
@@ -147,19 +147,6 @@ transcription is configured, slash commands, and `@` mentions for installed Apps
or MCP presets. The model badge shows the current model or preset and links back or MCP presets. The model badge shows the current model or preset and links back
to model settings when setup is incomplete. to model settings when setup is incomplete.
When two or more named model presets are configured, the badge shows a dropdown
indicator and acts as a preset selector. Click or tap it, then choose the preset
you want from the menu. For keyboard access, focus the badge and press
<kbd>Enter</kbd> or <kbd>Space</kbd> to open the menu, use the arrow keys to move,
and press <kbd>Enter</kbd> to select.
The selection applies to future turns in the current session and persists with
that session; it does not change the default for other sessions. Only named
presets from **Settings → Models** are selectable. An inline fallback model that
has not been saved as a named preset is not a separate manual choice. Save it as
a named preset to make it selectable. The same switch is available in chat with
`/model <preset>`; see [Chat Commands: Model Presets](./chat-commands.md#model-presets).
For image generation, configure an image provider first and then use the WebUI For image generation, configure an image provider first and then use the WebUI
image mode from the composer. See [`image-generation.md`](./image-generation.md) image mode from the composer. See [`image-generation.md`](./image-generation.md)
for provider setup and output behavior. for provider setup and output behavior.
+1 -4
View File
@@ -24,7 +24,6 @@ class ProviderSnapshot:
@dataclass(frozen=True) @dataclass(frozen=True)
class _ProviderSetup: class _ProviderSetup:
model: str model: str
provider_name: str
provider_config: ProviderConfig | None provider_config: ProviderConfig | None
spec: ProviderSpec | None spec: ProviderSpec | None
backend: str backend: str
@@ -100,7 +99,6 @@ def _resolve_provider_setup(
return _ProviderSetup( return _ProviderSetup(
model=model, model=model,
provider_name=provider_name,
provider_config=p, provider_config=p,
spec=spec, spec=spec,
backend=backend, backend=backend,
@@ -136,7 +134,6 @@ def _make_provider_core(
model=model, model=model,
) )
model = setup.model model = setup.model
provider_name = setup.provider_name
p = setup.provider_config p = setup.provider_config
spec = setup.spec spec = setup.spec
backend = setup.backend backend = setup.backend
@@ -201,7 +198,7 @@ def _make_provider_core(
extra_headers=_provider_extra_headers(spec, p), extra_headers=_provider_extra_headers(spec, p),
spec=spec, spec=spec,
extra_body=p.extra_body if p else None, extra_body=p.extra_body if p else None,
api_type=p.api_type if p and provider_name == "openai" else "auto", api_type=p.api_type if p else "auto",
extra_query=p.extra_query if p else None, extra_query=p.extra_query if p else None,
proxy=p.proxy if p else None, proxy=p.proxy if p else None,
) )
+46 -37
View File
@@ -49,7 +49,7 @@ from nanobot.providers.openai_responses import (
if TYPE_CHECKING: if TYPE_CHECKING:
from openai import AsyncOpenAI as AsyncOpenAIType from openai import AsyncOpenAI as AsyncOpenAIType
from nanobot.providers.registry import ProviderSpec from nanobot.providers.registry import ProviderSpec, ResponsesCapabilities
# Module-level placeholder — set lazily by _ensure_client on first real # Module-level placeholder — set lazily by _ensure_client on first real
# use, or replaced by tests via ``patch(...)``. Kept as a plain name so # use, or replaced by tests via ``patch(...)``. Kept as a plain name so
@@ -470,7 +470,12 @@ class OpenAICompatProvider(LLMProvider):
self.extra_headers = extra_headers or {} self.extra_headers = extra_headers or {}
self._spec = spec self._spec = spec
self._extra_body = extra_body or {} self._extra_body = extra_body or {}
self._api_type = api_type if spec and spec.name == "openai" else "auto" responses = spec.responses if spec is not None else None
self._api_type = (
api_type
if responses is not None and responses.allows_api_type_override
else "auto"
)
self._extra_query = extra_query or {} self._extra_query = extra_query or {}
self._proxy = proxy or None self._proxy = proxy or None
self._native_compaction_available = True self._native_compaction_available = True
@@ -961,39 +966,33 @@ class OpenAICompatProvider(LLMProvider):
"""Choose Responses for providers/models that explicitly support it.""" """Choose Responses for providers/models that explicitly support it."""
if self._api_type == "chat_completions": if self._api_type == "chat_completions":
return False return False
spec_name = self._spec.name if self._spec is not None else None capabilities = self._responses_capabilities()
model_name = self._request_model_name(model or self.default_model).lower() if capabilities is None:
supported_models = {
supported.lower()
for supported in getattr(self._spec, "responses_models", ())
}
model_responses = any(
model_name == supported or model_name.endswith(f"/{supported}")
for supported in supported_models
)
provider_responses = spec_name in ("openai", "github_copilot")
if not provider_responses and not model_responses:
return False return False
model_name = self._request_model_name(model or self.default_model).lower()
if self._api_type == "responses": if self._api_type == "responses":
# Explicit configuration means Responses is mandatory; do not # Explicit configuration means Responses is mandatory; do not
# consult the circuit breaker or fall back to Chat Completions. # consult the circuit breaker or fall back to Chat Completions.
return True return True
if provider_responses and (self._spec is None or self._spec.name != "github_copilot"): if (
if not _is_direct_openai_base(self._effective_base): capabilities.requires_direct_openai_base
and not _is_direct_openai_base(self._effective_base)
):
return False return False
wants = False explicitly_supported = capabilities.matches_model(model_name)
if model_responses: wants_auto_route = capabilities.auto_route and (
wants = True (reasoning_effort is not None and reasoning_effort.lower() != "none")
elif reasoning_effort and reasoning_effort.lower() != "none": or any(token in model_name for token in ("gpt-5", "o1", "o3", "o4"))
wants = True )
elif any(token in model_name for token in ("gpt-5", "o1", "o3", "o4")): if not explicitly_supported and not wants_auto_route:
wants = True
if not wants:
return False return False
return self._responses_circuit_allows_probe(model, reasoning_effort) return self._responses_circuit_allows_probe(model, reasoning_effort)
def _responses_capabilities(self) -> ResponsesCapabilities | None:
return self._spec.responses if self._spec is not None else None
def _responses_state_provider(self) -> str: def _responses_state_provider(self) -> str:
spec_name = self._spec.name if self._spec is not None else "custom" spec_name = self._spec.name if self._spec is not None else "custom"
effective_base = self._effective_base or "https://api.openai.com/v1" effective_base = self._effective_base or "https://api.openai.com/v1"
@@ -1016,14 +1015,20 @@ class OpenAICompatProvider(LLMProvider):
def supports_native_compaction(self, model: str | None = None) -> bool: def supports_native_compaction(self, model: str | None = None) -> bool:
"""Enable server compaction only on direct OpenAI Responses endpoints.""" """Enable server compaction only on direct OpenAI Responses endpoints."""
_ = model _ = model
capabilities = self._responses_capabilities()
if ( if (
not self._native_compaction_available not self._native_compaction_available
or self._api_type == "chat_completions" or self._api_type == "chat_completions"
or capabilities is None
or not capabilities.supports_native_compaction
): ):
return False return False
if self._spec is not None and self._spec.name != "openai": if (
capabilities.requires_direct_openai_base
and not _is_direct_openai_base(self._effective_base)
):
return False return False
return _is_direct_openai_base(self._effective_base) return True
def _responses_circuit_allows_probe( def _responses_circuit_allows_probe(
self, self,
@@ -1111,7 +1116,10 @@ class OpenAICompatProvider(LLMProvider):
self._sanitize_empty_content(sanitized_state.pending_messages) self._sanitize_empty_content(sanitized_state.pending_messages)
) )
) )
preserve_reasoning = bool(self._spec and self._spec.name == "deepseek") capabilities = self._responses_capabilities()
preserve_reasoning = (
capabilities is not None and capabilities.reasoning_replay == "plaintext"
)
instructions, input_items, replayed = prepare_responses_input( instructions, input_items, replayed = prepare_responses_input(
sanitized_messages, sanitized_messages,
state=sanitized_state, state=sanitized_state,
@@ -1142,10 +1150,15 @@ class OpenAICompatProvider(LLMProvider):
"compact_threshold": compact_threshold, "compact_threshold": compact_threshold,
}] }]
if self._supports_temperature(model_name, reasoning_effort): supports_temperature = self._supports_temperature(model_name, reasoning_effort)
if supports_temperature:
body["temperature"] = temperature body["temperature"] = temperature
if not self._supports_temperature(model_name, reasoning_effort) and not preserve_reasoning: if (
not supports_temperature
and capabilities is not None
and capabilities.reasoning_replay == "encrypted"
):
body["include"] = ["reasoning.encrypted_content"] body["include"] = ["reasoning.encrypted_content"]
if reasoning_effort and reasoning_effort.lower() != "none": if reasoning_effort and reasoning_effort.lower() != "none":
body["reasoning"] = {"effort": reasoning_effort} body["reasoning"] = {"effort": reasoning_effort}
@@ -1766,10 +1779,8 @@ class OpenAICompatProvider(LLMProvider):
self._record_responses_success(model, reasoning_effort) self._record_responses_success(model, reasoning_effort)
return result return result
except Exception as responses_error: except Exception as responses_error:
if self._spec and self._spec.name == "github_copilot": capabilities = self._responses_capabilities()
# Copilot gateway exposes GPT-5/o-series only via /responses; if capabilities is not None and not capabilities.allows_chat_fallback:
# falling back to /chat/completions cannot succeed and would
# hide the real error.
raise raise
if self._api_type == "responses": if self._api_type == "responses":
raise raise
@@ -1862,10 +1873,8 @@ class OpenAICompatProvider(LLMProvider):
) )
return result return result
except Exception as responses_error: except Exception as responses_error:
if self._spec and self._spec.name == "github_copilot": capabilities = self._responses_capabilities()
# Copilot gateway exposes GPT-5/o-series only via /responses; if capabilities is not None and not capabilities.allows_chat_fallback:
# falling back to /chat/completions cannot succeed and would
# hide the real error.
raise raise
if self._api_type == "responses": if self._api_type == "responses":
raise raise
+45 -6
View File
@@ -13,7 +13,7 @@ Every entry writes out all fields so you can copy-paste as a template.
from __future__ import annotations from __future__ import annotations
from dataclasses import dataclass from dataclasses import dataclass
from typing import Any from typing import Any, Literal
from pydantic.alias_generators import to_snake from pydantic.alias_generators import to_snake
@@ -28,6 +28,32 @@ class ProviderModelSpec:
context_window: int | None = None context_window: int | None = None
@dataclass(frozen=True)
class ResponsesCapabilities:
"""Provider capabilities for the shared OpenAI Responses execution path.
``reasoning_replay`` selects whether multi-turn reasoning is retained as
encrypted server content, plaintext local history, or not requested.
"""
models: tuple[str, ...] = ()
auto_route: bool = False
requires_direct_openai_base: bool = False
allows_api_type_override: bool = False
reasoning_replay: Literal["none", "encrypted", "plaintext"] = "none"
supports_native_compaction: bool = False
allows_chat_fallback: bool = True
def matches_model(self, model: str) -> bool:
"""Return whether *model* is explicitly routed through Responses."""
model_name = model.lower()
return any(
model_name == supported.lower()
or model_name.endswith(f"/{supported.lower()}")
for supported in self.models
)
@dataclass(frozen=True) @dataclass(frozen=True)
class ProviderSpec: class ProviderSpec:
"""One LLM provider's metadata. See PROVIDERS below for real examples. """One LLM provider's metadata. See PROVIDERS below for real examples.
@@ -111,10 +137,8 @@ class ProviderSpec:
# Substring match against the wire model name (lowercased). # Substring match against the wire model name (lowercased).
implicit_reasoning_models: tuple[str, ...] = () implicit_reasoning_models: tuple[str, ...] = ()
# Models that expose the OpenAI Responses wire format. This is model-level # Capabilities for providers/models served through the shared Responses path.
# because providers may add Responses support incrementally (DeepSeek V4 responses: ResponsesCapabilities | None = None
# Flash is supported before V4 Pro).
responses_models: tuple[str, ...] = ()
# When the model returns content as a list of {"type":"thinking",...} + # When the model returns content as a list of {"type":"thinking",...} +
# {"type":"text",...} blocks, extract the thinking text into # {"type":"text",...} blocks, extract the thinking text into
@@ -373,6 +397,13 @@ PROVIDERS: tuple[ProviderSpec, ...] = (
display_name="OpenAI", display_name="OpenAI",
backend="openai_compat", backend="openai_compat",
supports_max_completion_tokens=True, supports_max_completion_tokens=True,
responses=ResponsesCapabilities(
auto_route=True,
requires_direct_openai_base=True,
allows_api_type_override=True,
reasoning_replay="encrypted",
supports_native_compaction=True,
),
), ),
# OpenAI Codex: OAuth-based, dedicated provider # OpenAI Codex: OAuth-based, dedicated provider
ProviderSpec( ProviderSpec(
@@ -456,6 +487,11 @@ PROVIDERS: tuple[ProviderSpec, ...] = (
strip_model_prefix=True, strip_model_prefix=True,
is_oauth=True, is_oauth=True,
supports_max_completion_tokens=True, supports_max_completion_tokens=True,
responses=ResponsesCapabilities(
auto_route=True,
reasoning_replay="encrypted",
allows_chat_fallback=False,
),
), ),
# DeepSeek: OpenAI-compatible at api.deepseek.com # DeepSeek: OpenAI-compatible at api.deepseek.com
ProviderSpec( ProviderSpec(
@@ -466,7 +502,10 @@ PROVIDERS: tuple[ProviderSpec, ...] = (
backend="openai_compat", backend="openai_compat",
default_api_base="https://api.deepseek.com", default_api_base="https://api.deepseek.com",
thinking_style="thinking_type", thinking_style="thinking_type",
responses_models=("deepseek-v4-flash",), responses=ResponsesCapabilities(
models=("deepseek-v4-flash",),
reasoning_replay="plaintext",
),
), ),
# Gemini: Google's OpenAI-compatible endpoint # Gemini: Google's OpenAI-compatible endpoint
ProviderSpec( ProviderSpec(
@@ -48,6 +48,7 @@ def test_build_responses_body_strips_github_copilot_prefix():
provider_context=ProviderCallContext(context_window_tokens=128_000), provider_context=ProviderCallContext(context_window_tokens=128_000),
) )
assert body["model"] == "gpt-5.4-mini" assert body["model"] == "gpt-5.4-mini"
assert body["include"] == ["reasoning.encrypted_content"]
assert "context_management" not in body assert "context_management" not in body
@@ -10,6 +10,11 @@ from nanobot.providers.openai_compat_provider import (
_RESPONSES_PROBE_INTERVAL_S, _RESPONSES_PROBE_INTERVAL_S,
OpenAICompatProvider, OpenAICompatProvider,
) )
from nanobot.providers.registry import (
ProviderSpec,
ResponsesCapabilities,
find_by_name,
)
@pytest.fixture() @pytest.fixture()
@@ -17,7 +22,7 @@ def provider():
"""A direct-OpenAI provider with Responses API support.""" """A direct-OpenAI provider with Responses API support."""
p = OpenAICompatProvider.__new__(OpenAICompatProvider) p = OpenAICompatProvider.__new__(OpenAICompatProvider)
p.default_model = "gpt-5" p.default_model = "gpt-5"
p._spec = type("Spec", (), {"name": "openai"})() p._spec = find_by_name("openai")
p._effective_base = "https://api.openai.com/v1" p._effective_base = "https://api.openai.com/v1"
p._api_type = "auto" p._api_type = "auto"
p._responses_failures = {} p._responses_failures = {}
@@ -30,12 +35,7 @@ def test_responses_api_available_by_default(provider):
def test_deepseek_v4_flash_uses_responses_by_model(provider): def test_deepseek_v4_flash_uses_responses_by_model(provider):
provider._spec = type("Spec", (), { provider._spec = find_by_name("deepseek")
"name": "deepseek",
"responses_models": ("deepseek-v4-flash",),
"strip_model_prefix": False,
"strip_model_prefixes": (),
})()
provider._effective_base = "https://api.deepseek.com" provider._effective_base = "https://api.deepseek.com"
provider.default_model = "deepseek-v4-flash" provider.default_model = "deepseek-v4-flash"
@@ -44,17 +44,48 @@ def test_deepseek_v4_flash_uses_responses_by_model(provider):
def test_deepseek_v4_flash_matches_provider_prefixed_model(provider): def test_deepseek_v4_flash_matches_provider_prefixed_model(provider):
provider._spec = type("Spec", (), { provider._spec = find_by_name("deepseek")
"name": "deepseek",
"responses_models": ("deepseek-v4-flash",),
"strip_model_prefix": False,
"strip_model_prefixes": (),
})()
provider._effective_base = "https://api.deepseek.com" provider._effective_base = "https://api.deepseek.com"
assert provider._should_use_responses_api("deepseek/deepseek-v4-flash", None) is True assert provider._should_use_responses_api("deepseek/deepseek-v4-flash", None) is True
def test_responses_behavior_is_declared_by_capabilities(provider):
provider._spec = ProviderSpec(
name="example",
keywords=("example",),
env_key="EXAMPLE_API_KEY",
responses=ResponsesCapabilities(
models=("example-o3",),
reasoning_replay="plaintext",
),
)
provider._effective_base = "https://example.test"
assert provider._should_use_responses_api("example-o3", None) is True
body = provider._build_responses_body(
messages=[
{"role": "user", "content": "question"},
{
"role": "assistant",
"reasoning_content": "think first",
"content": "answer",
},
{"role": "user", "content": "follow-up"},
],
tools=None,
model="example-o3",
max_tokens=100,
temperature=0.1,
reasoning_effort="high",
tool_choice=None,
)
assert {"type": "reasoning", "content": "think first"} in body["input"]
assert "include" not in body
def test_direct_openai_enables_server_compaction(provider): def test_direct_openai_enables_server_compaction(provider):
provider._extra_body = {} provider._extra_body = {}
@@ -73,6 +104,7 @@ def test_direct_openai_enables_server_compaction(provider):
"type": "compaction", "type": "compaction",
"compact_threshold": 70_000, "compact_threshold": 70_000,
}] }]
assert body["include"] == ["reasoning.encrypted_content"]
def test_api_type_chat_completions_disables_responses(provider): def test_api_type_chat_completions_disables_responses(provider):
@@ -96,7 +128,7 @@ def test_api_type_responses_ignores_circuit_breaker(provider):
def test_api_type_responses_does_not_force_non_openai(provider): def test_api_type_responses_does_not_force_non_openai(provider):
provider._spec = type("Spec", (), {"name": "custom"})() provider._spec = find_by_name("custom")
provider._api_type = "responses" provider._api_type = "responses"
assert provider._should_use_responses_api("gpt-4o", None) is False assert provider._should_use_responses_api("gpt-4o", None) is False
+250 -86
View File
@@ -1,13 +1,13 @@
import { useLayoutEffect, useRef, useState } from "react";
import { ChevronDown, CircleHelp, Sparkles } from "lucide-react";
import { import {
DropdownMenu, useEffect,
DropdownMenuContent, useLayoutEffect,
DropdownMenuRadioGroup, useRef,
DropdownMenuRadioItem, useState,
DropdownMenuTrigger, type KeyboardEvent,
} from "@/components/ui/dropdown-menu"; type PointerEvent,
} from "react";
import { CircleHelp, Sparkles } from "lucide-react";
import { useLogoFallback } from "@/hooks/useLogoFallback"; import { useLogoFallback } from "@/hooks/useLogoFallback";
import { inferProviderFromModelName, providerBrand } from "@/lib/provider-brand"; import { inferProviderFromModelName, providerBrand } from "@/lib/provider-brand";
import { cn } from "@/lib/utils"; import { cn } from "@/lib/utils";
@@ -33,6 +33,54 @@ interface ModelPresetBadgeProps {
onClick?: () => void; onClick?: () => void;
} }
interface PresetGesture {
active: boolean;
baseIndex: number;
latestY: number;
pointerId: number;
startY: number;
step: number;
target: HTMLElement;
timer: ReturnType<typeof setTimeout> | null;
}
interface PresetMotion {
index: number;
remainder: number;
settling: boolean;
}
const LONG_PRESS_MS = 400;
const PRESS_SLOP_PX = 8;
const PILL_GAP_PX = 4;
const PILL_OFFSETS = [-2, -1, 0, 1, 2] as const;
const HANDOFF_THRESHOLD = 0.56;
const DOCK_MAX_SCALE = 1.08;
const DOCK_RADIUS = 1.5;
const SETTLE_MS = 180;
function wrapIndex(index: number, length: number): number {
return ((index % length) + length) % length;
}
function dockScale(distanceFromFocus: number): number {
const distance = Math.abs(distanceFromFocus);
if (distance >= DOCK_RADIUS) return 1;
const influence = (1 + Math.cos(Math.PI * distance / DOCK_RADIUS)) / 2;
return 1 + (DOCK_MAX_SCALE - 1) * influence;
}
function stepWithHysteresis(raw: number, current: number): number {
let next = current;
while (raw > next + HANDOFF_THRESHOLD) next += 1;
while (raw < next - HANDOFF_THRESHOLD) next -= 1;
return next;
}
function preventTouchScroll(event: TouchEvent) {
if (event.cancelable) event.preventDefault();
}
export function ModelPresetBadge({ export function ModelPresetBadge({
label, label,
modelDetail, modelDetail,
@@ -62,13 +110,150 @@ export function ModelPresetBadge({
: modelPresets.map((preset, index) => index === listedIndex ? activePreset : preset); : modelPresets.map((preset, index) => index === listedIndex ? activePreset : preset);
const interactive = Boolean(onClick); const interactive = Boolean(onClick);
const canSwitch = !interactive && Boolean(onPresetChange) && activeName !== "" && presets.length > 1; const canSwitch = !interactive && Boolean(onPresetChange) && activeName !== "" && presets.length > 1;
const badgeClassName = cn( const currentIndex = Math.max(0, presets.findIndex((preset) => preset.name === activeName));
const pillHeight = isHero ? 32 : 36;
const pillStride = pillHeight + PILL_GAP_PX;
const [motion, setMotion] = useState<PresetMotion | null>(null);
const gestureRef = useRef<PresetGesture | null>(null);
function clearGesture() {
const gesture = gestureRef.current;
if (gesture?.timer) clearTimeout(gesture.timer);
if (gesture?.active) gesture.target.removeEventListener("touchmove", preventTouchScroll);
gestureRef.current = null;
}
useEffect(() => {
if (!canSwitch) {
clearGesture();
setMotion(null);
}
return clearGesture;
}, [canSwitch]);
useEffect(() => {
if (!motion?.settling) return;
const timer = setTimeout(() => setMotion(null), SETTLE_MS + 80);
return () => clearTimeout(timer);
}, [motion?.settling]);
function updateMotion(gesture: PresetGesture, clientY: number) {
const raw = -(clientY - gesture.startY) / pillStride;
gesture.step = stepWithHysteresis(raw, gesture.step);
setMotion({ index: gesture.baseIndex + gesture.step, remainder: raw - gesture.step, settling: false });
}
function handlePointerDown(event: PointerEvent<HTMLElement>) {
if (!canSwitch || gestureRef.current || motion || event.isPrimary === false) return;
if (event.pointerType === "mouse" && event.button !== 0) return;
const gesture: PresetGesture = {
active: false,
baseIndex: currentIndex,
latestY: event.clientY,
pointerId: event.pointerId,
startY: event.clientY,
step: 0,
target: event.currentTarget,
timer: null,
};
gesture.timer = setTimeout(() => {
if (gestureRef.current !== gesture) return;
gesture.active = true;
updateMotion(gesture, gesture.latestY);
gesture.target.addEventListener("touchmove", preventTouchScroll, { passive: false });
try {
gesture.target.setPointerCapture(gesture.pointerId);
} catch { /* The pointer may already have ended. */ }
}, LONG_PRESS_MS);
gestureRef.current = gesture;
}
function handlePointerMove(event: PointerEvent<HTMLElement>) {
const gesture = gestureRef.current;
if (!gesture || gesture.pointerId !== event.pointerId) return;
gesture.latestY = event.clientY;
if (!gesture.active) {
if (Math.abs(event.clientY - gesture.startY) > PRESS_SLOP_PX) clearGesture();
return;
}
event.preventDefault();
updateMotion(gesture, event.clientY);
}
function finishGesture(event: PointerEvent<HTMLElement>, commit: boolean) {
const gesture = gestureRef.current;
if (!gesture || gesture.pointerId !== event.pointerId) return;
clearGesture();
if (event.currentTarget.hasPointerCapture?.(gesture.pointerId)) {
event.currentTarget.releasePointerCapture?.(gesture.pointerId);
}
if (!commit || !gesture.active) {
setMotion(null);
return;
}
const selected = presets[wrapIndex(gesture.baseIndex + gesture.step, presets.length)];
setMotion((current) => current && { ...current, remainder: 0, settling: true });
if (selected && selected.name !== activeName) onPresetChange?.(selected.name);
}
function handleKeyDown(event: KeyboardEvent<HTMLElement>) {
if (!canSwitch) return;
const targetByKey: Record<string, number> = {
ArrowUp: currentIndex - 1,
ArrowDown: currentIndex + 1,
Home: 0,
End: presets.length - 1,
};
const target = targetByKey[event.key];
if (target === undefined) return;
event.preventDefault();
const next = presets[wrapIndex(target, presets.length)];
if (next?.name !== activeName) onPresetChange?.(next.name);
}
const previewIndex = wrapIndex(motion?.index ?? currentIndex, presets.length);
const previewPreset = presets[previewIndex];
const Container = interactive || canSwitch ? "button" : "span";
const trackOffset = motion ? -pillStride * (2 + motion.remainder) : 0;
return (
<Container
data-switching={motion ? "true" : undefined}
data-settling={motion?.settling ? "true" : undefined}
aria-label={label}
aria-orientation={canSwitch ? "vertical" : undefined}
aria-valuemax={canSwitch ? presets.length - 1 : undefined}
aria-valuemin={canSwitch ? 0 : undefined}
aria-valuenow={canSwitch ? previewIndex : undefined}
aria-valuetext={canSwitch ? previewPreset?.label || label : undefined}
role={canSwitch ? "spinbutton" : undefined}
type={interactive || canSwitch ? "button" : undefined}
onClick={interactive ? onClick : undefined}
onKeyDown={handleKeyDown}
onPointerDown={handlePointerDown}
onPointerMove={handlePointerMove}
onPointerLeave={(event) => {
const gesture = gestureRef.current;
if (gesture && gesture.pointerId === event.pointerId && !gesture.active) clearGesture();
}}
onPointerUp={(event) => finishGesture(event, true)}
onPointerCancel={(event) => finishGesture(event, false)}
onLostPointerCapture={(event) => finishGesture(event, false)}
onContextMenu={(event) => {
if (gestureRef.current?.active) event.preventDefault();
}}
onDragStart={(event) => event.preventDefault()}
style={{ touchAction: canSwitch ? "manipulation" : undefined }}
className={cn(
"thread-composer-model-badge group/model-badge relative inline-flex w-fit min-w-0 max-w-[min(18rem,44vw)] justify-end appearance-none border-0 bg-transparent p-0 shadow-none", "thread-composer-model-badge group/model-badge relative inline-flex w-fit min-w-0 max-w-[min(18rem,44vw)] justify-end appearance-none border-0 bg-transparent p-0 shadow-none",
(interactive || canSwitch) && "cursor-pointer focus-visible:outline-none", interactive && "cursor-pointer",
canSwitch && "cursor-grab select-none focus-visible:outline-none",
motion && "z-10 cursor-grabbing",
isHero ? "h-8" : "h-9", isHero ? "h-8" : "h-9",
); )}
const badgeContent = ( >
<PresetPill <PresetPill
className={motion && "invisible"}
label={label} label={label}
modelDetail={modelDetail} modelDetail={modelDetail}
provider={provider} provider={provider}
@@ -76,80 +261,53 @@ export function ModelPresetBadge({
needsSetup={needsSetup} needsSetup={needsSetup}
fallbackModelName={fallbackModelName} fallbackModelName={fallbackModelName}
isHero={isHero} isHero={isHero}
showPicker={canSwitch}
/> />
); {motion ? (
<span
if (canSwitch) { data-testid="composer-model-pill-viewport"
return ( className={cn(
<DropdownMenu modal={false}> "composer-model-pill-viewport pointer-events-none absolute right-0 w-max max-w-[calc(44vw+0.5rem)] overflow-hidden bg-transparent pl-2 sm:max-w-[18.5rem]",
<DropdownMenuTrigger asChild> isHero ? "-bottom-2.5 -top-2.5" : "-bottom-3 -top-3",
<button type="button" aria-label={label} className={badgeClassName}> )}
{badgeContent} aria-hidden
</button>
</DropdownMenuTrigger>
<DropdownMenuContent
align="end"
side="top"
sideOffset={8}
collisionPadding={12}
className="w-[min(20rem,calc(100vw-2rem))] rounded-[18px]"
> >
<DropdownMenuRadioGroup <span
value={activeName} data-testid="composer-model-pill-track"
onValueChange={(name) => { data-settling={motion.settling ? "true" : undefined}
if (name !== activeName) onPresetChange?.(name); className="composer-model-pill-track ml-auto flex w-max max-w-full flex-col items-end gap-1 will-change-transform"
onTransitionEnd={(event) => {
if (motion.settling && event.currentTarget === event.target) setMotion(null);
}}
style={{
paddingTop: isHero ? "10px" : "12px",
transform: `translate3d(0, ${trackOffset}px, 0)`,
}} }}
> >
{presets.map((preset) => { {PILL_OFFSETS.map((offset) => {
const detail = [...new Set([preset.model, preset.provider].filter(Boolean))] const virtualIndex = motion.index + offset;
.join(" · "); const preset = presets[wrapIndex(virtualIndex, presets.length)];
const scale = motion.settling ? 1 : dockScale(offset - motion.remainder);
return ( return (
<DropdownMenuRadioItem <PresetPill
key={preset.name} key={virtualIndex}
value={preset.name} label={preset.label || preset.name}
className="min-h-[46px] items-start rounded-[14px] py-2.5" modelDetail={preset.model}
> provider={preset.provider}
<span className="min-w-0 flex-1"> isHero={isHero}
<span className="block truncate font-semibold text-foreground"> offset={offset}
{preset.label || preset.name} scale={scale}
</span> />
{detail ? (
<span className="mt-0.5 block truncate text-[11.5px] text-muted-foreground">
{detail}
</span>
) : null}
</span>
</DropdownMenuRadioItem>
); );
})} })}
</DropdownMenuRadioGroup>
</DropdownMenuContent>
</DropdownMenu>
);
}
if (interactive) {
return (
<button
type="button"
aria-label={label}
onClick={onClick}
className={badgeClassName}
>
{badgeContent}
</button>
);
}
return (
<span aria-label={label} className={badgeClassName}>
{badgeContent}
</span> </span>
</span>
) : null}
</Container>
); );
} }
function PresetPill({ function PresetPill({
className,
label, label,
modelDetail, modelDetail,
provider, provider,
@@ -157,8 +315,10 @@ function PresetPill({
needsSetup = false, needsSetup = false,
fallbackModelName, fallbackModelName,
isHero, isHero,
showPicker = false, offset,
scale,
}: { }: {
className?: string | false | null;
label: string; label: string;
modelDetail?: string | null; modelDetail?: string | null;
provider?: string | null; provider?: string | null;
@@ -166,7 +326,8 @@ function PresetPill({
needsSetup?: boolean; needsSetup?: boolean;
fallbackModelName?: string | null; fallbackModelName?: string | null;
isHero: boolean; isHero: boolean;
showPicker?: boolean; offset?: number;
scale?: number;
}) { }) {
const labelRef = useRef<HTMLSpanElement | null>(null); const labelRef = useRef<HTMLSpanElement | null>(null);
const [labelOverflows, setLabelOverflows] = useState(false); const [labelOverflows, setLabelOverflows] = useState(false);
@@ -176,7 +337,9 @@ function PresetPill({
const brand = providerBrand(inferredProvider); const brand = providerBrand(inferredProvider);
const { logoUrl, onLogoError, onLogoLoad } = useLogoFallback(brand?.logoUrls); const { logoUrl, onLogoError, onLogoLoad } = useLogoFallback(brand?.logoUrls);
const title = [...new Set([label, modelDetail, providerLabel].filter(Boolean))].join(" · "); const title = [...new Set([label, modelDetail, providerLabel].filter(Boolean))].join(" · ");
const logoTestId = needsSetup const logoTestId = offset !== undefined
? undefined
: needsSetup
? "composer-model-setup-icon" ? "composer-model-setup-icon"
: `composer-model-logo${inferredProvider ? `-${inferredProvider}` : ""}`; : `composer-model-logo${inferredProvider ? `-${inferredProvider}` : ""}`;
@@ -193,15 +356,22 @@ function PresetPill({
return ( return (
<span <span
data-fallback={fallbackModelName ? "true" : undefined} data-fallback={fallbackModelName ? "true" : undefined}
data-preset-offset={offset}
title={fallbackModelName || title || undefined} title={fallbackModelName || title || undefined}
className={cn( className={cn(
"composer-model-badge composer-model-pill inline-flex h-full w-fit max-w-full min-w-0 shrink-0 items-center rounded-full border border-border/55 bg-card font-medium text-foreground/70", "composer-model-badge composer-model-pill inline-flex h-full w-fit max-w-full min-w-0 shrink-0 items-center rounded-full border border-border/55 bg-card font-medium text-foreground/70",
"shadow-[0_2px_8px_rgba(15,23,42,0.045)]", offset === undefined && "shadow-[0_2px_8px_rgba(15,23,42,0.045)]",
"transition-[color,background-color,border-color,transform] duration-150 ease-out group-focus-visible/model-badge:ring-2 group-focus-visible/model-badge:ring-ring/45", "transition-[color,background-color,border-color,transform] duration-150 ease-out group-focus-visible/model-badge:ring-2 group-focus-visible/model-badge:ring-ring/45",
showPicker && "group-hover/model-badge:border-border group-hover/model-badge:text-foreground/85",
needsSetup && "border-amber-500/35 bg-amber-50/70 text-amber-900 dark:bg-amber-500/10 dark:text-amber-200", needsSetup && "border-amber-500/35 bg-amber-50/70 text-amber-900 dark:bg-amber-500/10 dark:text-amber-200",
isHero ? "gap-1.5 px-2.5 text-[12px]" : "gap-2 px-3 text-[12.5px]", isHero ? "gap-1.5 px-2.5 text-[12px]" : "gap-2 px-3 text-[12.5px]",
offset !== undefined && "composer-model-pill-dock",
className,
)} )}
style={scale === undefined ? undefined : {
height: `${isHero ? 32 : 36}px`,
transform: `scale(${scale.toFixed(4)})`,
zIndex: Math.round(scale * 100),
}}
> >
<span <span
data-testid={logoTestId} data-testid={logoTestId}
@@ -252,12 +422,6 @@ function PresetPill({
> >
{label} {label}
</span> </span>
{showPicker ? (
<ChevronDown
className="thread-composer-model-chevron h-3.5 w-3.5 shrink-0 text-muted-foreground/75"
aria-hidden
/>
) : null}
</span> </span>
); );
} }
+41 -5
View File
@@ -738,14 +738,54 @@
mask-image: linear-gradient(to right, #000 0, #000 calc(100% - 0.75rem), transparent); mask-image: linear-gradient(to right, #000 0, #000 calc(100% - 0.75rem), transparent);
} }
.thread-composer-model-badge:active > .composer-model-pill { .thread-composer-model-badge:not([data-switching="true"]):active
> .composer-model-pill {
transform: scale(0.98); transform: scale(0.98);
} }
@keyframes composer-model-pill-viewport-enter {
from {
transform: scale(0.9074);
}
to {
transform: scale(1);
}
}
.composer-model-pill-viewport {
transform-origin: right center;
animation: composer-model-pill-viewport-enter 210ms
cubic-bezier(0.2, 0.8, 0.2, 1) both;
-webkit-mask-image: linear-gradient(to bottom, transparent, #000 4px, #000 calc(100% - 4px), transparent);
mask-image: linear-gradient(to bottom, transparent, #000 4px, #000 calc(100% - 4px), transparent);
}
.composer-model-pill-dock {
transform-origin: right center;
transition-property: none;
will-change: transform;
}
.composer-model-pill-track[data-settling="true"],
.composer-model-pill-track[data-settling="true"] .composer-model-pill-dock {
transition: transform 180ms cubic-bezier(0.22, 1, 0.36, 1);
}
@media (prefers-reduced-motion: reduce) { @media (prefers-reduced-motion: reduce) {
.thread-composer-model-badge:active > .composer-model-pill { .thread-composer-model-badge:active > .composer-model-pill {
transform: none !important; transform: none !important;
} }
.composer-model-pill-track[data-settling="true"],
.composer-model-pill-dock {
transition: none;
will-change: auto;
}
.composer-model-pill-viewport {
animation: none;
}
} }
@container thread-composer (max-width: 21rem) { @container thread-composer (max-width: 21rem) {
@@ -798,10 +838,6 @@
.thread-composer-model-label { .thread-composer-model-label {
display: none; display: none;
} }
.thread-composer-model-chevron {
display: none;
}
} }
@container thread-composer (max-width: 16rem) { @container thread-composer (max-width: 16rem) {
+98 -21
View File
@@ -1,5 +1,4 @@
import { act, fireEvent, render, screen, waitFor, within } from "@testing-library/react"; import { act, fireEvent, render, screen, waitFor, within } from "@testing-library/react";
import userEvent from "@testing-library/user-event";
import { afterEach, describe, expect, it, vi } from "vitest"; import { afterEach, describe, expect, it, vi } from "vitest";
import { ThreadComposer } from "@/components/thread/ThreadComposer"; import { ThreadComposer } from "@/components/thread/ThreadComposer";
@@ -314,11 +313,28 @@ function renderPresetComposer(variant: "thread" | "hero" = "thread") {
/>, />,
); );
return { return {
badge: screen.getByRole("button", { name: "Kimi" }), badge: screen.getByRole("spinbutton", { name: "Kimi" }),
onPresetChange, onPresetChange,
}; };
} }
function pointerDown(badge: HTMLElement, pointerId = 7, clientY = 100, button = 0) {
fireEvent.pointerDown(badge, {
button,
clientY,
isPrimary: true,
pointerId,
pointerType: "mouse",
});
}
function longPress(badge: HTMLElement, pointerId = 7) {
pointerDown(badge, pointerId);
act(() => {
vi.advanceTimersByTime(400);
});
}
describe("ThreadComposer", () => { describe("ThreadComposer", () => {
it("focuses and sends a removable quoted answer excerpt", async () => { it("focuses and sends a removable quoted answer excerpt", async () => {
const onSend = vi.fn(); const onSend = vi.fn();
@@ -412,7 +428,7 @@ describe("ThreadComposer", () => {
/>, />,
); );
const badge = screen.getByRole("button", { name: "gpt-5.6-sol" }); const badge = screen.getByRole("spinbutton", { name: "gpt-5.6-sol" });
expect(badge).toHaveClass("w-fit", "max-w-[min(18rem,44vw)]"); expect(badge).toHaveClass("w-fit", "max-w-[min(18rem,44vw)]");
expect(badge).not.toHaveClass("w-[5.75rem]"); expect(badge).not.toHaveClass("w-[5.75rem]");
expect(screen.getByText("gpt-5.6-sol")).toBeInTheDocument(); expect(screen.getByText("gpt-5.6-sol")).toBeInTheDocument();
@@ -445,32 +461,93 @@ describe("ThreadComposer", () => {
expect(screen.queryByText(/Enter to send/)).not.toBeInTheDocument(); expect(screen.queryByText(/Enter to send/)).not.toBeInTheDocument();
}); });
it("opens a preset menu on click and switches the selected preset", async () => { it("scrolls complete preset pills after a left-button long press and wraps", () => {
const user = userEvent.setup(); vi.useFakeTimers();
const { badge, onPresetChange } = renderPresetComposer(); const { badge, onPresetChange } = renderPresetComposer();
expect(badge).toHaveClass("h-9"); expect(badge).toHaveClass("h-9");
expect(badge).toHaveAttribute("aria-haspopup", "menu"); expect(badge).toHaveStyle({ touchAction: "manipulation" });
expect(badge).toHaveAttribute("aria-expanded", "false"); const idleTouchMove = new Event("touchmove", {
bubbles: true,
cancelable: true,
});
badge.dispatchEvent(idleTouchMove);
expect(idleTouchMove.defaultPrevented).toBe(false);
fireEvent.click(badge);
pointerDown(badge);
fireEvent.pointerMove(badge, { clientY: 80, pointerId: 7, pointerType: "mouse" });
act(() => vi.advanceTimersByTime(500));
fireEvent.pointerUp(badge, { clientY: 80, pointerId: 7, pointerType: "mouse" });
expect(onPresetChange).not.toHaveBeenCalled();
await user.click(badge); longPress(badge);
expect(badge).toHaveAttribute("data-switching", "true");
const viewport = screen.getByTestId("composer-model-pill-viewport");
expect(viewport).toHaveClass(
"right-0",
"w-max",
"max-w-[calc(44vw+0.5rem)]",
"overflow-hidden",
"-top-3",
"-bottom-3",
);
const track = screen.getByTestId("composer-model-pill-track");
expect(track).toHaveClass("w-max", "max-w-full", "items-end", "gap-1");
const activeTouchMove = new Event("touchmove", {
bubbles: true,
cancelable: true,
});
badge.dispatchEvent(activeTouchMove);
expect(activeTouchMove.defaultPrevented).toBe(true);
const pills = track.querySelectorAll<HTMLElement>(".composer-model-pill");
expect(pills).toHaveLength(5);
expect(Array.from(pills).every((pill) => pill.classList.contains("w-fit"))).toBe(true);
expect(Array.from(pills).every((pill) => pill.querySelector("img"))).toBe(true);
expect(Array.from(badge.querySelectorAll("img")).every((image) => !image.draggable)).toBe(true);
const centeredPill = track.querySelector<HTMLElement>("[data-preset-offset='0']");
expect(centeredPill).toHaveTextContent("Kimi");
expect(centeredPill).toHaveStyle({ transform: "scale(1.0800)" });
expect(
track.querySelector<HTMLElement>("[data-preset-offset='1']"),
).toHaveStyle({ transform: "scale(1.0200)" });
expect(badge).toHaveAttribute("aria-expanded", "true"); fireEvent.pointerMove(badge, {
expect(screen.getByRole("menuitemradio", { name: /Kimi.*moonshot/i })) clientY: 122,
.toHaveAttribute("aria-checked", "true"); pointerId: 7,
expect(screen.getByRole("menuitemradio", { name: /DFlash.*deepseek/i })) pointerType: "mouse",
.toBeInTheDocument(); });
await user.click(screen.getByRole("menuitemradio", { name: /DS Pro.*deepseek/i })); expect(track.querySelector("[data-preset-offset='0']")).toHaveTextContent("Kimi");
expect(onPresetChange).toHaveBeenCalledWith("dspro"); fireEvent.pointerMove(badge, {
expect(screen.queryByRole("menu")).not.toBeInTheDocument(); clientY: 123,
pointerId: 7,
pointerType: "mouse",
});
expect(track.querySelector("[data-preset-offset='0']")).toHaveTextContent("DS Pro");
fireEvent.pointerUp(badge, {
clientY: 123,
pointerId: 7,
pointerType: "mouse",
}); });
it("supports the same preset menu in hero mode", async () => { expect(onPresetChange).toHaveBeenCalledWith("dspro");
const user = userEvent.setup(); expect(badge).toHaveAttribute("data-settling", "true");
expect(track).toHaveAttribute("data-settling", "true");
act(() => {
vi.advanceTimersByTime(260);
});
expect(badge).not.toHaveAttribute("data-switching");
expect(badge).not.toHaveAttribute("data-settling");
});
it("supports the same long-press switcher in hero mode and cancels pointercancel", () => {
vi.useFakeTimers();
const { badge, onPresetChange } = renderPresetComposer("hero"); const { badge, onPresetChange } = renderPresetComposer("hero");
expect(badge).toHaveClass("h-8"); expect(badge).toHaveClass("h-8");
await user.click(badge); longPress(badge, 9);
await user.click(screen.getByRole("menuitemradio", { name: /DFlash.*deepseek/i })); expect(badge).toHaveAttribute("data-switching", "true");
expect(onPresetChange).toHaveBeenCalledWith("dflash"); fireEvent.pointerMove(badge, { clientY: 75, pointerId: 9, pointerType: "mouse" });
fireEvent.pointerCancel(badge, { clientY: 75, pointerId: 9, pointerType: "mouse" });
expect(badge).not.toHaveAttribute("data-switching");
expect(onPresetChange).not.toHaveBeenCalled();
}); });
it("transcribes voice input into the composer without sending", async () => { it("transcribes voice input into the composer without sending", async () => {
+10 -7
View File
@@ -586,18 +586,19 @@ describe("ThreadShell", () => {
)); ));
const { rerender } = render(view("default")); const { rerender } = render(view("default"));
const badge = await screen.findByRole("button", { name: "Default" }); const badge = await screen.findByRole("spinbutton", { name: "Default" });
expect(badge).toHaveTextContent("Default"); expect(badge).toHaveTextContent("Default");
fireEvent.pointerDown(badge); fireEvent.keyDown(badge, { key: "ArrowDown" });
fireEvent.click(await screen.findByRole("menuitemradio", { name: /^Fast/ }));
expect(client.sendSystemCommand).toHaveBeenCalledWith( expect(client.sendSystemCommand).toHaveBeenCalledWith(
"preset-order", "preset-order",
"/model fast", "/model fast",
); );
expect(await screen.findByText("Fast")).toBeInTheDocument(); expect(await screen.findByText("Fast")).toBeInTheDocument();
fireEvent.pointerDown(screen.getByRole("button", { name: "Fast" })); fireEvent.keyDown(
fireEvent.click(await screen.findByRole("menuitemradio", { name: /^Extra/ })); screen.getByRole("spinbutton", { name: "Fast" }),
{ key: "End" },
);
expect(client.sendSystemCommand).toHaveBeenLastCalledWith( expect(client.sendSystemCommand).toHaveBeenLastCalledWith(
"preset-order", "preset-order",
"/model extra", "/model extra",
@@ -971,8 +972,10 @@ describe("ThreadShell", () => {
)); ));
const { rerender } = render(view(null)); const { rerender } = render(view(null));
fireEvent.pointerDown(await screen.findByRole("button", { name: "Default" })); fireEvent.keyDown(
fireEvent.click(await screen.findByRole("menuitemradio", { name: /^Fast/ })); await screen.findByRole("spinbutton", { name: "Default" }),
{ key: "ArrowDown" },
);
expect(await screen.findByText("Fast")).toBeInTheDocument(); expect(await screen.findByText("Fast")).toBeInTheDocument();
expect(client.sendSystemCommand).not.toHaveBeenCalled(); expect(client.sendSystemCommand).not.toHaveBeenCalled();