fix(models): synchronize canonical runtime selection

This commit is contained in:
Xubin Ren
2026-08-16 11:50:56 +08:00
parent 0a6ee1c539
commit 731b8fc2ed
15 changed files with 265 additions and 16 deletions
+9 -3
View File
@@ -533,9 +533,15 @@ class AgentLoop:
self.subagents.max_iterations = self.max_iterations self.subagents.max_iterations = self.max_iterations
def invalidate_runtime_config(self) -> None: def invalidate_runtime_config(self) -> None:
"""Invalidate runtime config and notify clients to refresh its catalog.""" """Invalidate runtime config for lazy refresh at the next admission."""
self.runtime_resolver.invalidate() self.runtime_resolver.invalidate()
self._publish_runtime_selection(self.runtime_resolver.runtime)
def refresh_runtime_config(self) -> LLMRuntime:
"""Refresh runtime config now and publish the canonical selection."""
self.runtime_resolver.invalidate()
runtime = self.runtime_resolver.admit()
self._publish_runtime_selection(runtime)
return runtime
def runtime_for_session( def runtime_for_session(
self, self,
@@ -1787,7 +1793,7 @@ class AgentLoop:
session.provider_state = None session.provider_state = None
self.sessions.save(session) self.sessions.save(session)
ctx.input_persisted_early = True ctx.input_persisted_early = True
ctx.delivery.record_runtime(runtime) await ctx.delivery.runtime_admitted(runtime)
ctx.request_context = self._request_context_for_turn(ctx) ctx.request_context = self._request_context_for_turn(ctx)
if ctx.kind is TurnKind.USER: if ctx.kind is TurnKind.USER:
+9 -1
View File
@@ -189,7 +189,15 @@ class TurnDelivery:
started_at=started_at, started_at=started_at,
) )
def record_runtime(self, runtime: LLMRuntime) -> None: async def runtime_admitted(self, runtime: LLMRuntime) -> None:
"""Record the immutable runtime and expose it at the lifecycle seam."""
if self.route.publish_lifecycle:
await self.runtime_event_publisher.turn_runtime_admitted(
self.delivery_message,
self.session_key,
runtime,
)
return
self.runtime_event_publisher.record_turn_runtime(self.session_key, runtime) self.runtime_event_publisher.record_turn_runtime(self.session_key, runtime)
def record_latency(self, latency_ms: int | None) -> None: def record_latency(self, latency_ms: int | None) -> None:
+2 -1
View File
@@ -84,9 +84,10 @@ class RuntimeModelUpdatedEvent(OutboundEvent):
@dataclass(frozen=True) @dataclass(frozen=True)
class TurnModelUpdatedEvent(OutboundEvent): class TurnModelUpdatedEvent(OutboundEvent):
"""The fallback model currently handling one chat turn.""" """The canonical preset and concrete model handling one chat turn."""
model: str model: str
model_preset: str | None = None
def outbound_message_for_event( def outbound_message_for_event(
+30
View File
@@ -40,6 +40,14 @@ class SessionTurnStarted:
context: RuntimeEventContext context: RuntimeEventContext
@dataclass(frozen=True)
class TurnRuntimeAdmitted:
"""The immutable model runtime selected for one admitted turn."""
context: RuntimeEventContext
runtime: LLMRuntime
@dataclass(frozen=True) @dataclass(frozen=True)
class TurnRunStatusChanged: class TurnRunStatusChanged:
"""Visible run status changed for a turn.""" """Visible run status changed for a turn."""
@@ -85,6 +93,7 @@ class RuntimeModelChanged:
RuntimeEvent = ( RuntimeEvent = (
SessionTurnStarted SessionTurnStarted
| TurnRuntimeAdmitted
| SessionTurnPersisted | SessionTurnPersisted
| TurnRunStatusChanged | TurnRunStatusChanged
| TurnCompleted | TurnCompleted
@@ -93,6 +102,7 @@ RuntimeEvent = (
) )
RuntimeEventType = ( RuntimeEventType = (
type[SessionTurnStarted] type[SessionTurnStarted]
| type[TurnRuntimeAdmitted]
| type[SessionTurnPersisted] | type[SessionTurnPersisted]
| type[TurnRunStatusChanged] | type[TurnRunStatusChanged]
| type[TurnCompleted] | type[TurnCompleted]
@@ -204,6 +214,26 @@ class RuntimeEventPublisher:
) )
) )
async def turn_runtime_admitted(
self,
msg: InboundMessage,
session_key: str,
runtime: LLMRuntime,
) -> None:
"""Record and publish the runtime selected for one turn."""
self.record_turn_runtime(session_key, runtime)
await self.bus.publish(
TurnRuntimeAdmitted(
context=self._context(
channel=msg.channel,
chat_id=msg.chat_id,
session_key=session_key,
metadata=msg.metadata,
),
runtime=runtime,
)
)
async def run_status_changed( async def run_status_changed(
self, self,
msg: InboundMessage, msg: InboundMessage,
+4
View File
@@ -1407,6 +1407,7 @@ class WebSocketChannel(BaseChannel):
await self.send_turn_model_updated( await self.send_turn_model_updated(
msg.chat_id, msg.chat_id,
model_name=event.model, model_name=event.model,
model_preset=event.model_preset,
) )
return return
if isinstance(event, GoalStateSyncEvent): if isinstance(event, GoalStateSyncEvent):
@@ -1774,6 +1775,7 @@ class WebSocketChannel(BaseChannel):
chat_id: str, chat_id: str,
*, *,
model_name: Any, model_name: Any,
model_preset: Any = None,
) -> None: ) -> None:
"""Notify one chat's subscribers which model is handling its current request.""" """Notify one chat's subscribers which model is handling its current request."""
conns = list(self._subs.get(chat_id, ())) conns = list(self._subs.get(chat_id, ()))
@@ -1788,6 +1790,8 @@ class WebSocketChannel(BaseChannel):
"chat_id": chat_id, "chat_id": chat_id,
"model_name": model_name.strip(), "model_name": model_name.strip(),
} }
if isinstance(model_preset, str) and model_preset.strip():
body["model_preset"] = model_preset.strip()
raw = json.dumps(body, ensure_ascii=False) raw = json.dumps(body, ensure_ascii=False)
for connection in conns: for connection in conns:
await self._safe_send_to(connection, raw, label=" turn_model_updated ") await self._safe_send_to(connection, raw, label=" turn_model_updated ")
@@ -1642,7 +1642,10 @@ async def test_send_scopes_turn_model_updates_to_the_subscribed_chat() -> None:
channel="websocket", channel="websocket",
chat_id="chat-1", chat_id="chat-1",
content="", content="",
event=TurnModelUpdatedEvent(model="deepseek/deepseek-chat"), event=TurnModelUpdatedEvent(
model="deepseek/deepseek-chat",
model_preset="Deep Research",
),
) )
) )
@@ -1651,6 +1654,7 @@ async def test_send_scopes_turn_model_updates_to_the_subscribed_chat() -> None:
"event": "turn_model_updated", "event": "turn_model_updated",
"chat_id": "chat-1", "chat_id": "chat-1",
"model_name": "deepseek/deepseek-chat", "model_name": "deepseek/deepseek-chat",
"model_preset": "Deep Research",
} }
chat_two.send.assert_not_awaited() chat_two.send.assert_not_awaited()
+1 -1
View File
@@ -658,7 +658,7 @@ def _run_gateway(
return agent.model.strip() or None return agent.model.strip() or None
def _webui_refresh_runtime_config() -> None: def _webui_refresh_runtime_config() -> None:
agent.invalidate_runtime_config() agent.refresh_runtime_config()
def _webui_skill_state_action(disabled_skills: set[str]) -> None: def _webui_skill_state_action(disabled_skills: set[str]) -> None:
config.agents.defaults.disabled_skills = sorted(disabled_skills) config.agents.defaults.disabled_skills = sorted(disabled_skills)
+28 -1
View File
@@ -33,6 +33,7 @@ from nanobot.bus.runtime_events import (
SessionTurnStarted, SessionTurnStarted,
TurnCompleted, TurnCompleted,
TurnRunStatusChanged, TurnRunStatusChanged,
TurnRuntimeAdmitted,
) )
from nanobot.providers.base import LLMProvider from nanobot.providers.base import LLMProvider
from nanobot.providers.fallback_provider import FallbackModelObserver from nanobot.providers.fallback_provider import FallbackModelObserver
@@ -459,7 +460,14 @@ def build_webui_fallback_model_observer(bus: MessageBus) -> FallbackModelObserve
outbound_message_for_event( outbound_message_for_event(
channel=context.channel, channel=context.channel,
chat_id=chat_id, chat_id=chat_id,
event=TurnModelUpdatedEvent(model=model), event=TurnModelUpdatedEvent(
model=model,
model_preset=(
context.runtime.model_preset
if context.runtime is not None
else None
),
),
metadata=context.metadata, metadata=context.metadata,
) )
) )
@@ -486,6 +494,10 @@ class WebuiTurnCoordinator:
self._handle_run_status_changed, self._handle_run_status_changed,
TurnRunStatusChanged, TurnRunStatusChanged,
), ),
runtime_events.subscribe(
self._handle_turn_runtime_admitted,
TurnRuntimeAdmitted,
),
runtime_events.subscribe( runtime_events.subscribe(
self._handle_turn_completed_event, self._handle_turn_completed_event,
TurnCompleted, TurnCompleted,
@@ -537,6 +549,21 @@ class WebuiTurnCoordinator:
started_at=event.started_at, started_at=event.started_at,
) )
async def _handle_turn_runtime_admitted(self, event: TurnRuntimeAdmitted) -> None:
if not self._is_websocket_event(event.context):
return
await self.bus.publish_outbound(
outbound_message_for_event(
channel=event.context.channel,
chat_id=event.context.chat_id,
event=TurnModelUpdatedEvent(
model=event.runtime.model,
model_preset=event.runtime.model_preset,
),
metadata=event.context.metadata,
)
)
async def _handle_turn_completed_event(self, event: TurnCompleted) -> None: async def _handle_turn_completed_event(self, event: TurnCompleted) -> None:
if not self._is_websocket_event(event.context): if not self._is_websocket_event(event.context):
return return
+9 -2
View File
@@ -1646,6 +1646,11 @@ class ModelSettingsHandler:
self.settings = settings self.settings = settings
self.logger = logger self.logger = logger
def _refresh_runtime_config(self) -> None:
"""Make a successful model-settings mutation visible to live clients now."""
if self.settings.refresh_runtime_config is not None:
self.settings.refresh_runtime_config()
async def handle( async def handle(
self, self,
action: str, action: str,
@@ -1655,6 +1660,7 @@ class ModelSettingsHandler:
try: try:
if action == "agent-update": if action == "agent-update":
payload = self.settings.mutate(operations.update_agent, request.query) payload = self.settings.mutate(operations.update_agent, request.query)
self._refresh_runtime_config()
return SettingsRouteResult.success( return SettingsRouteResult.success(
payload, payload,
decorate_restart=True, decorate_restart=True,
@@ -1667,8 +1673,7 @@ class ModelSettingsHandler:
request.query, request.query,
rename_model_preset=self.settings.rename_model_preset, rename_model_preset=self.settings.rename_model_preset,
) )
if self.settings.refresh_runtime_config is not None: self._refresh_runtime_config()
self.settings.refresh_runtime_config()
return SettingsRouteResult.success(payload, decorate_restart=True) return SettingsRouteResult.success(payload, decorate_restart=True)
mutation = { mutation = {
@@ -1680,6 +1685,7 @@ class ModelSettingsHandler:
}.get(action) }.get(action)
if mutation is not None: if mutation is not None:
payload = self.settings.mutate(mutation, request.query) payload = self.settings.mutate(mutation, request.query)
self._refresh_runtime_config()
return SettingsRouteResult.success(payload, decorate_restart=True) return SettingsRouteResult.success(payload, decorate_restart=True)
if action == "provider-update": if action == "provider-update":
@@ -1690,6 +1696,7 @@ class ModelSettingsHandler:
payload, image_restart_cleared = await operations.apply_image_runtime_change( payload, image_restart_cleared = await operations.apply_image_runtime_change(
payload payload
) )
self._refresh_runtime_config()
return SettingsRouteResult.success( return SettingsRouteResult.success(
payload, payload,
decorate_restart=True, decorate_restart=True,
+51 -4
View File
@@ -226,7 +226,7 @@ def test_named_default_refresh_is_used_by_sessions_without_override(tmp_path: Pa
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_config_invalidation_notifies_clients_before_session_runtime_refresh( async def test_config_invalidation_defers_canonical_notification_until_default_refresh(
tmp_path: Path, tmp_path: Path,
) -> None: ) -> None:
provider = _provider("model-a") provider = _provider("model-a")
@@ -266,12 +266,59 @@ async def test_config_invalidation_notifies_clients_before_session_runtime_refre
runtime = loop.runtime_for_session(session) runtime = loop.runtime_for_session(session)
await asyncio.sleep(0) await asyncio.sleep(0)
assert [(event.model, event.model_preset) for event in published] == [ assert published == []
("model-a", "fast"),
]
assert runtime.model == "model-b" assert runtime.model == "model-b"
assert loop.model_presets["fast"].model == "model-b" assert loop.model_presets["fast"].model == "model-b"
assert loop.llm_runtime().model == "model-b"
await asyncio.sleep(0)
assert [(event.model, event.model_preset) for event in published] == [
("model-b", "fast"),
]
@pytest.mark.asyncio
async def test_config_refresh_publishes_renamed_canonical_preset(tmp_path: Path) -> None:
provider = _provider("model-a")
catalog = {"fast": ModelPresetConfig(model="model-a")}
default_name = "fast"
published: list[RuntimeModelChanged] = []
def load_preset(name: str) -> ProviderSnapshot:
return ProviderSnapshot(
provider=provider,
model=catalog[name].model,
context_window_tokens=16_000,
signature=(name, catalog[name].model),
model_preset=name,
)
loop = AgentLoop(
bus=MessageBus(),
provider=provider,
workspace=tmp_path,
model="model-a",
context_window_tokens=16_000,
provider_signature=("fast", "model-a"),
provider_snapshot_loader=lambda: load_preset(default_name),
model_presets=catalog,
preset_catalog_loader=lambda: catalog,
model_preset="fast",
preset_snapshot_loader=load_preset,
)
loop.runtime_events.subscribe(published.append, RuntimeModelChanged)
catalog["Codex"] = catalog.pop("fast")
default_name = "Codex"
runtime = loop.refresh_runtime_config()
await asyncio.sleep(0)
assert (runtime.model, runtime.model_preset) == ("model-a", "Codex")
assert [(event.model, event.model_preset) for event in published] == [
("model-a", "Codex"),
]
def test_next_turn_captures_generation_changed_after_previous_admission( def test_next_turn_captures_generation_changed_after_previous_admission(
tmp_path: Path, tmp_path: Path,
+32
View File
@@ -10,6 +10,7 @@ from nanobot.bus.runtime_events import (
SessionTurnStarted, SessionTurnStarted,
TurnCompleted, TurnCompleted,
TurnRunStatusChanged, TurnRunStatusChanged,
TurnRuntimeAdmitted,
) )
@@ -123,6 +124,37 @@ async def test_runtime_event_publisher_consumes_turn_metadata_on_complete() -> N
assert second.runtime is None assert second.runtime is None
@pytest.mark.asyncio
async def test_runtime_event_publisher_exposes_admitted_runtime() -> None:
bus = RuntimeEventBus()
seen: list[object] = []
publisher = RuntimeEventPublisher(bus)
msg = InboundMessage(
channel="websocket",
sender_id="user",
chat_id="chat-a",
content="hello",
)
runtime = object()
bus.subscribe(seen.append)
await publisher.turn_runtime_admitted(msg, "websocket:chat-a", runtime) # type: ignore[arg-type]
await publisher.turn_completed(
channel="websocket",
chat_id="chat-a",
session_key="websocket:chat-a",
metadata=None,
)
admitted = seen[0]
completed = seen[1]
assert isinstance(admitted, TurnRuntimeAdmitted)
assert admitted.runtime is runtime
assert admitted.context.chat_id == "chat-a"
assert isinstance(completed, TurnCompleted)
assert completed.runtime is runtime
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_runtime_event_publisher_emits_persisted_turn_attributes() -> None: async def test_runtime_event_publisher_emits_persisted_turn_attributes() -> None:
bus = RuntimeEventBus() bus = RuntimeEventBus()
+52
View File
@@ -7,7 +7,11 @@ import pytest
from nanobot.agent.tools.context import RequestContext, request_context from nanobot.agent.tools.context import RequestContext, request_context
from nanobot.bus.events import InboundMessage from nanobot.bus.events import InboundMessage
from nanobot.bus.outbound_events import GoalStatusEvent, TurnModelUpdatedEvent from nanobot.bus.outbound_events import GoalStatusEvent, TurnModelUpdatedEvent
from nanobot.bus.runtime_events import RuntimeEventBus, RuntimeEventContext, TurnRuntimeAdmitted
from nanobot.providers.base import GenerationSettings
from nanobot.session import webui_turns as wth from nanobot.session import webui_turns as wth
from nanobot.session.manager import SessionManager
from nanobot.utils.llm_runtime import LLMRuntime
from nanobot.webui.metadata import WEBSOCKET_TURN_OWNER_METADATA_KEY from nanobot.webui.metadata import WEBSOCKET_TURN_OWNER_METADATA_KEY
@@ -144,10 +148,18 @@ async def test_fallback_model_is_scoped_to_its_websocket_chat() -> None:
bus.publish_outbound = AsyncMock() bus.publish_outbound = AsyncMock()
observer = wth.build_webui_fallback_model_observer(bus) observer = wth.build_webui_fallback_model_observer(bus)
runtime = LLMRuntime(
provider=MagicMock(),
model="openai/gpt-4.1",
generation=GenerationSettings(),
context_window_tokens=16_000,
model_preset="Deep Research",
)
with request_context( with request_context(
RequestContext( RequestContext(
channel="websocket", channel="websocket",
chat_id="chat-model", chat_id="chat-model",
runtime=runtime,
metadata={"webui": True}, metadata={"webui": True},
) )
): ):
@@ -159,6 +171,46 @@ async def test_fallback_model_is_scoped_to_its_websocket_chat() -> None:
assert outbound.metadata == {"webui": True} assert outbound.metadata == {"webui": True}
assert isinstance(outbound.event, TurnModelUpdatedEvent) assert isinstance(outbound.event, TurnModelUpdatedEvent)
assert outbound.event.model == "deepseek/deepseek-chat" assert outbound.event.model == "deepseek/deepseek-chat"
assert outbound.event.model_preset == "Deep Research"
@pytest.mark.asyncio
async def test_admitted_runtime_publishes_chat_scoped_model_and_preset(tmp_path) -> None:
bus = MagicMock()
bus.publish_outbound = AsyncMock()
runtime_events = RuntimeEventBus()
coordinator = wth.WebuiTurnCoordinator(
bus=bus,
sessions=SessionManager(tmp_path),
schedule_background=lambda coro: coro.close(),
)
coordinator.subscribe(runtime_events)
runtime = LLMRuntime(
provider=MagicMock(),
model="openai-codex/gpt-5.6",
generation=GenerationSettings(),
context_window_tokens=262_144,
model_preset="Codex",
)
await runtime_events.publish(
TurnRuntimeAdmitted(
context=RuntimeEventContext(
channel="websocket",
chat_id="chat-model",
session_key="websocket:chat-model",
metadata={"webui": True},
),
runtime=runtime,
)
)
outbound = bus.publish_outbound.await_args.args[0]
assert outbound.channel == "websocket"
assert outbound.chat_id == "chat-model"
assert isinstance(outbound.event, TurnModelUpdatedEvent)
assert outbound.event.model == "openai-codex/gpt-5.6"
assert outbound.event.model_preset == "Codex"
@pytest.mark.asyncio @pytest.mark.asyncio
+30 -2
View File
@@ -307,6 +307,18 @@ async def test_oauth_completion_reads_websocket_payload(
@pytest.mark.parametrize( @pytest.mark.parametrize(
("route_path", "function_name", "payload", "expected_query"), ("route_path", "function_name", "payload", "expected_query"),
[ [
(
"/api/settings/update",
"update_agent_settings",
{"model_preset": "Codex"},
{"model_preset": ["Codex"]},
),
(
"/api/settings/model-configurations/create",
"create_model_configuration",
{"name": "Codex", "model": "openai-codex/gpt-5.6"},
{"name": ["Codex"], "model": ["openai-codex/gpt-5.6"]},
),
( (
"/api/settings/model-configurations/delete", "/api/settings/model-configurations/delete",
"delete_model_configuration", "delete_model_configuration",
@@ -325,10 +337,22 @@ async def test_oauth_completion_reads_websocket_payload(
{"order": ["backup"]}, {"order": ["backup"]},
{"order": ['["backup"]']}, {"order": ['["backup"]']},
), ),
(
"/api/settings/provider/create",
"create_provider_settings",
{"name": "team", "api_base": "https://llm.example/v1"},
{"name": ["team"], "api_base": ["https://llm.example/v1"]},
),
(
"/api/settings/provider/update",
"update_provider_settings",
{"provider": "team", "api_base": "https://llm.example/v2"},
{"provider": ["team"], "api_base": ["https://llm.example/v2"]},
),
], ],
) )
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_model_preset_mutation_routes( async def test_runtime_config_mutation_routes_refresh_live_runtime(
monkeypatch, monkeypatch,
route_path: str, route_path: str,
function_name: str, function_name: str,
@@ -336,6 +360,7 @@ async def test_model_preset_mutation_routes(
expected_query: dict[str, list[str]], expected_query: dict[str, list[str]],
) -> None: ) -> None:
captured: dict[str, object] = {} captured: dict[str, object] = {}
refresh_runtime_config = MagicMock()
def mutate(query, *, config_path=None): def mutate(query, *, config_path=None):
captured["query"] = query captured["query"] = query
@@ -344,12 +369,15 @@ async def test_model_preset_mutation_routes(
monkeypatch.setattr(f"nanobot.webui.settings_routes.{function_name}", mutate) monkeypatch.setattr(f"nanobot.webui.settings_routes.{function_name}", mutate)
request = _mutation_request(route_path, payload) request = _mutation_request(route_path, payload)
response = await _router().dispatch(None, request, route_path) response = await _router(
refresh_runtime_config=refresh_runtime_config,
).dispatch(None, request, route_path)
assert response is not None assert response is not None
assert response.status_code == 200 assert response.status_code == 200
assert json.loads(response.body)["routed"] == function_name assert json.loads(response.body)["routed"] == function_name
assert captured["query"] == expected_query assert captured["query"] == expected_query
refresh_runtime_config.assert_called_once_with()
@pytest.mark.asyncio @pytest.mark.asyncio
+1
View File
@@ -1271,6 +1271,7 @@ export type InboundEvent =
event: "turn_model_updated"; event: "turn_model_updated";
chat_id: string; chat_id: string;
model_name: string; model_name: string;
model_preset?: string | null;
} }
| ({ | ({
event: "turn_end"; event: "turn_end";
+2
View File
@@ -1636,12 +1636,14 @@ describe("NanobotClient", () => {
event: "turn_model_updated", event: "turn_model_updated",
chat_id: "chat-a", chat_id: "chat-a",
model_name: "deepseek/deepseek-chat", model_name: "deepseek/deepseek-chat",
model_preset: "Deep Research",
}); });
expect(chatHandler).toHaveBeenCalledWith({ expect(chatHandler).toHaveBeenCalledWith({
event: "turn_model_updated", event: "turn_model_updated",
chat_id: "chat-a", chat_id: "chat-a",
model_name: "deepseek/deepseek-chat", model_name: "deepseek/deepseek-chat",
model_preset: "Deep Research",
}); });
}); });