refactor(agent): move provider refresh into subsystem owners

This commit is contained in:
Xubin Ren
2026-04-26 14:18:37 +00:00
parent f670da6c70
commit b2aec5528a
4 changed files with 25 additions and 11 deletions
+3 -10
View File
@@ -307,16 +307,9 @@ class AgentLoop:
self.model = model
self.context_window_tokens = context_window_tokens
self.runner.provider = provider
self.subagents.provider = provider
self.subagents.model = model
self.subagents.runner.provider = provider
self.consolidator.provider = provider
self.consolidator.model = model
self.consolidator.context_window_tokens = context_window_tokens
self.consolidator.max_completion_tokens = provider.generation.max_tokens
self.dream.provider = provider
self.dream.model = model
self.dream._runner.provider = provider
self.subagents.set_provider(provider, model)
self.consolidator.set_provider(provider, model, context_window_tokens)
self.dream.set_provider(provider, model)
self._provider_signature = snapshot.signature
logger.info("Runtime model switched for next turn: {} -> {}", old_model, model)
+16
View File
@@ -450,6 +450,17 @@ class Consolidator:
weakref.WeakValueDictionary()
)
def set_provider(
self,
provider: LLMProvider,
model: str,
context_window_tokens: int,
) -> None:
self.provider = provider
self.model = model
self.context_window_tokens = context_window_tokens
self.max_completion_tokens = provider.generation.max_tokens
def get_lock(self, session_key: str) -> asyncio.Lock:
"""Return the shared consolidation lock for one session."""
return self._locks.setdefault(session_key, asyncio.Lock())
@@ -710,6 +721,11 @@ class Dream:
self._runner = AgentRunner(provider)
self._tools = self._build_tools()
def set_provider(self, provider: LLMProvider, model: str) -> None:
self.provider = provider
self.model = model
self._runner.provider = provider
# -- tool registry -------------------------------------------------------
def _build_tools(self) -> ToolRegistry:
+5
View File
@@ -96,6 +96,11 @@ class SubagentManager:
self._task_statuses: dict[str, SubagentStatus] = {}
self._session_tasks: dict[str, set[str]] = {} # session_key -> {task_id, ...}
def set_provider(self, provider: LLMProvider, model: str) -> None:
self.provider = provider
self.model = model
self.runner.provider = provider
async def spawn(
self,
task: str,
+1 -1
View File
@@ -14,7 +14,7 @@ def _provider(default_model: str, max_tokens: int = 123) -> MagicMock:
return provider
def test_runtime_refresh_updates_loop_dependents(tmp_path: Path) -> None:
def test_provider_refresh_updates_all_model_dependents(tmp_path: Path) -> None:
old_provider = _provider("old-model")
new_provider = _provider("new-model", max_tokens=456)
loop = AgentLoop(