refactor(agent): move provider refresh into subsystem owners
This commit is contained in:
+3
-10
@@ -307,16 +307,9 @@ class AgentLoop:
|
|||||||
self.model = model
|
self.model = model
|
||||||
self.context_window_tokens = context_window_tokens
|
self.context_window_tokens = context_window_tokens
|
||||||
self.runner.provider = provider
|
self.runner.provider = provider
|
||||||
self.subagents.provider = provider
|
self.subagents.set_provider(provider, model)
|
||||||
self.subagents.model = model
|
self.consolidator.set_provider(provider, model, context_window_tokens)
|
||||||
self.subagents.runner.provider = provider
|
self.dream.set_provider(provider, model)
|
||||||
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._provider_signature = snapshot.signature
|
self._provider_signature = snapshot.signature
|
||||||
logger.info("Runtime model switched for next turn: {} -> {}", old_model, model)
|
logger.info("Runtime model switched for next turn: {} -> {}", old_model, model)
|
||||||
|
|
||||||
|
|||||||
@@ -450,6 +450,17 @@ class Consolidator:
|
|||||||
weakref.WeakValueDictionary()
|
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:
|
def get_lock(self, session_key: str) -> asyncio.Lock:
|
||||||
"""Return the shared consolidation lock for one session."""
|
"""Return the shared consolidation lock for one session."""
|
||||||
return self._locks.setdefault(session_key, asyncio.Lock())
|
return self._locks.setdefault(session_key, asyncio.Lock())
|
||||||
@@ -710,6 +721,11 @@ class Dream:
|
|||||||
self._runner = AgentRunner(provider)
|
self._runner = AgentRunner(provider)
|
||||||
self._tools = self._build_tools()
|
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 -------------------------------------------------------
|
# -- tool registry -------------------------------------------------------
|
||||||
|
|
||||||
def _build_tools(self) -> ToolRegistry:
|
def _build_tools(self) -> ToolRegistry:
|
||||||
|
|||||||
@@ -96,6 +96,11 @@ class SubagentManager:
|
|||||||
self._task_statuses: dict[str, SubagentStatus] = {}
|
self._task_statuses: dict[str, SubagentStatus] = {}
|
||||||
self._session_tasks: dict[str, set[str]] = {} # session_key -> {task_id, ...}
|
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(
|
async def spawn(
|
||||||
self,
|
self,
|
||||||
task: str,
|
task: str,
|
||||||
|
|||||||
@@ -14,7 +14,7 @@ def _provider(default_model: str, max_tokens: int = 123) -> MagicMock:
|
|||||||
return provider
|
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")
|
old_provider = _provider("old-model")
|
||||||
new_provider = _provider("new-model", max_tokens=456)
|
new_provider = _provider("new-model", max_tokens=456)
|
||||||
loop = AgentLoop(
|
loop = AgentLoop(
|
||||||
|
|||||||
Reference in New Issue
Block a user