mirror of
https://github.com/galaxyproject/galaxy.git
synced 2026-09-24 16:30:27 +08:00
Type the agents stack for the context it actually receives
The agents stack declared trans as ProvidesUserContext everywhere, so reaching the toolbox panel views needed a cast back to SessionRequestContext. Nothing was ever narrower: DependsOnTrans is declared to return a SessionRequestContext, and the MCP entry point builds one explicitly, so the declarations were widening what callers hand over. Declare SessionRequestContext along that chain instead, from the endpoints through AgentService and GalaxyAgentDependencies to AgentOperationsManager, and drop the cast. This also makes the toolbox search attribute resolve, so its type: ignore goes too.
This commit is contained in:
@@ -29,13 +29,13 @@ from typing import (
|
||||
import yaml
|
||||
|
||||
from galaxy.exceptions import ConfigurationError
|
||||
from galaxy.managers.context import ProvidesUserContext
|
||||
from galaxy.model import User
|
||||
from galaxy.schema.agents import (
|
||||
ActionSuggestion,
|
||||
ActionType,
|
||||
ConfidenceLevel,
|
||||
)
|
||||
from galaxy.work.context import SessionRequestContext
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from galaxy.config import GalaxyAppConfiguration
|
||||
@@ -353,7 +353,7 @@ class AgentRunState:
|
||||
class GalaxyAgentDependencies:
|
||||
"""Dependencies passed to Galaxy agents via dependency injection."""
|
||||
|
||||
trans: ProvidesUserContext
|
||||
trans: SessionRequestContext
|
||||
user: User
|
||||
config: "GalaxyAppConfiguration"
|
||||
# Callable to get agent instances, avoids circular import in base.py
|
||||
|
||||
@@ -7,14 +7,12 @@ Delegates to the Galaxy service layer for validation, permission checks, and pag
|
||||
import logging
|
||||
from typing import (
|
||||
Any,
|
||||
cast,
|
||||
Literal,
|
||||
)
|
||||
|
||||
from sqlalchemy import select
|
||||
|
||||
from galaxy.agents import iwc
|
||||
from galaxy.managers.context import ProvidesUserContext
|
||||
from galaxy.managers.hdas import HDAManager
|
||||
from galaxy.managers.tools import DynamicToolManager
|
||||
from galaxy.model import UserDynamicToolAssociation
|
||||
@@ -63,7 +61,7 @@ ID_FIELDS = {
|
||||
class AgentOperationsManager:
|
||||
"""Shared operations for AI agents, delegating to Galaxy's service layer."""
|
||||
|
||||
def __init__(self, app: MinimalManagerApp, trans: ProvidesUserContext):
|
||||
def __init__(self, app: MinimalManagerApp, trans: SessionRequestContext):
|
||||
self.app = app
|
||||
self.trans = trans
|
||||
self._tools_service: Any | None = None
|
||||
@@ -636,13 +634,10 @@ class AgentOperationsManager:
|
||||
}
|
||||
|
||||
def get_tool_panel(self, view: str | None = None) -> dict[str, Any]:
|
||||
# The agents stack types trans as ProvidesUserContext throughout, but at runtime it is
|
||||
# always the SessionRequestContext the panel-view methods require.
|
||||
trans = cast("SessionRequestContext", self.trans)
|
||||
if view is None:
|
||||
view = self.app.toolbox._default_panel_view(trans)
|
||||
view = self.app.toolbox._default_panel_view(self.trans)
|
||||
|
||||
tool_panel = self.app.toolbox.to_panel_view(trans, view=view)
|
||||
tool_panel = self.app.toolbox.to_panel_view(self.trans, view=view)
|
||||
|
||||
return {"tool_panel": tool_panel, "view": view}
|
||||
|
||||
|
||||
@@ -246,7 +246,7 @@ class ToolRecommendationAgent(BaseGalaxyAgent):
|
||||
|
||||
try:
|
||||
panel_view = self.deps.config.default_panel_view or "default"
|
||||
toolbox_search = self.deps.trans.app.toolbox_search # type: ignore[attr-defined]
|
||||
toolbox_search = self.deps.trans.app.toolbox_search
|
||||
tool_ids = toolbox_search.search(query, panel_view, self.deps.config)
|
||||
|
||||
tools = []
|
||||
|
||||
@@ -9,10 +9,10 @@ from galaxy.agents import GalaxyAgentDependencies
|
||||
from galaxy.agents.registry import AgentRegistry
|
||||
from galaxy.agents.router import QueryRouterAgent
|
||||
from galaxy.config import GalaxyAppConfiguration
|
||||
from galaxy.managers.context import ProvidesUserContext
|
||||
from galaxy.managers.jobs import JobManager
|
||||
from galaxy.model import User
|
||||
from galaxy.schema.agents import AgentResponse
|
||||
from galaxy.work.context import SessionRequestContext
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
@@ -30,7 +30,7 @@ class AgentService:
|
||||
self.job_manager = job_manager
|
||||
self.registry = registry
|
||||
|
||||
def create_dependencies(self, trans: ProvidesUserContext, user: User) -> GalaxyAgentDependencies:
|
||||
def create_dependencies(self, trans: SessionRequestContext, user: User) -> GalaxyAgentDependencies:
|
||||
"""Create agent dependencies for dependency injection."""
|
||||
toolbox = trans.app.toolbox if hasattr(trans, "app") and hasattr(trans.app, "toolbox") else None
|
||||
return GalaxyAgentDependencies(
|
||||
@@ -47,7 +47,7 @@ class AgentService:
|
||||
self,
|
||||
agent_type: str,
|
||||
query: str,
|
||||
trans: ProvidesUserContext,
|
||||
trans: SessionRequestContext,
|
||||
user: User,
|
||||
context: dict[str, Any] | None = None,
|
||||
) -> AgentResponse:
|
||||
@@ -96,7 +96,7 @@ class AgentService:
|
||||
async def route_and_execute(
|
||||
self,
|
||||
query: str,
|
||||
trans: ProvidesUserContext,
|
||||
trans: SessionRequestContext,
|
||||
user: User,
|
||||
context: dict[str, Any] | None = None,
|
||||
agent_type: str = "auto",
|
||||
|
||||
@@ -30,6 +30,7 @@ from galaxy.webapps.galaxy.api import (
|
||||
DependsOnUser,
|
||||
Router,
|
||||
)
|
||||
from galaxy.work.context import SessionRequestContext
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
@@ -90,7 +91,7 @@ class AgentAPI:
|
||||
async def query_agent(
|
||||
self,
|
||||
request: AgentQueryRequest,
|
||||
trans: ProvidesUserContext = DependsOnTrans,
|
||||
trans: SessionRequestContext = DependsOnTrans,
|
||||
user: User = DependsOnUser,
|
||||
) -> AgentQueryResponse:
|
||||
"""Query an AI agent. Use agent_type='auto' for automatic routing.
|
||||
@@ -129,7 +130,7 @@ class AgentAPI:
|
||||
job_id: DecodedDatabaseIdField | None = Body(None, description="Job ID for context"),
|
||||
error_details: dict[str, Any] | None = Body(None, description="Additional error details"),
|
||||
save_exchange: bool | None = Body(None, description="Save exchange for feedback tracking. Defaults to false."),
|
||||
trans: ProvidesUserContext = DependsOnTrans,
|
||||
trans: SessionRequestContext = DependsOnTrans,
|
||||
user: User = DependsOnUser,
|
||||
) -> AgentResponse:
|
||||
"""Analyze job errors and provide debugging assistance.
|
||||
@@ -181,7 +182,7 @@ class AgentAPI:
|
||||
query: str = Body(..., description="Description of the tool to create"),
|
||||
context: dict[str, Any] | None = Body(None, description="Additional context for tool creation"),
|
||||
save_exchange: bool | None = Body(None, description="Save exchange for feedback tracking. Defaults to false."),
|
||||
trans: ProvidesUserContext = DependsOnTrans,
|
||||
trans: SessionRequestContext = DependsOnTrans,
|
||||
user: User = DependsOnUser,
|
||||
) -> AgentResponse:
|
||||
"""Create a custom Galaxy tool.
|
||||
@@ -216,7 +217,7 @@ class AgentAPI:
|
||||
async def history_summary(
|
||||
self,
|
||||
history_id: str = Body(..., embed=True, description="Encoded id of the history to summarize."),
|
||||
trans: ProvidesUserContext = DependsOnTrans,
|
||||
trans: SessionRequestContext = DependsOnTrans,
|
||||
user: User = DependsOnUser,
|
||||
) -> AgentResponse:
|
||||
"""Produce a comprehensive markdown report for a history's analysis.
|
||||
|
||||
@@ -28,7 +28,6 @@ from galaxy.exceptions import (
|
||||
from galaxy.managers.agents import AgentService
|
||||
from galaxy.managers.chat import ChatManager
|
||||
from galaxy.managers.context import (
|
||||
ProvidesHistoryContext,
|
||||
ProvidesUserContext,
|
||||
)
|
||||
from galaxy.managers.jobs import JobManager
|
||||
@@ -53,6 +52,7 @@ from galaxy.webapps.galaxy.api import (
|
||||
DependsOnUser,
|
||||
Router,
|
||||
)
|
||||
from galaxy.work.context import SessionRequestContext
|
||||
|
||||
# Import agent system
|
||||
try:
|
||||
@@ -132,7 +132,7 @@ class ChatAPI:
|
||||
payload: ChatPayload | None = None,
|
||||
query: str | None = Query(default=None, description="Query string for general chat"),
|
||||
agent_type: str = Query(default="auto", description="Agent type to use for the query"),
|
||||
trans: ProvidesHistoryContext = DependsOnTrans,
|
||||
trans: SessionRequestContext = DependsOnTrans,
|
||||
user: User = DependsOnUser,
|
||||
) -> ChatResponse:
|
||||
"""GalaxyAI endpoint - handles both job-based and general chat queries
|
||||
@@ -412,7 +412,7 @@ class ChatAPI:
|
||||
workflow_id: str = Path(..., description="Workflow ID to generate the report for"),
|
||||
version: int | None = Query(None, description="Version of the workflow"),
|
||||
instance: bool = Query(False, description="Whether the workflow_id is an instance ID"),
|
||||
trans: ProvidesUserContext = DependsOnTrans,
|
||||
trans: SessionRequestContext = DependsOnTrans,
|
||||
user: User = DependsOnUser,
|
||||
) -> WorkflowReportResponse:
|
||||
"""Generate a report for the specified workflow."""
|
||||
@@ -566,7 +566,7 @@ class ChatAPI:
|
||||
self,
|
||||
query: str,
|
||||
agent_type: str,
|
||||
trans: ProvidesUserContext,
|
||||
trans: SessionRequestContext,
|
||||
user: User,
|
||||
job=None,
|
||||
context: dict[str, Any] | None = None,
|
||||
@@ -579,7 +579,7 @@ class ChatAPI:
|
||||
self,
|
||||
query: str,
|
||||
agent_type: str,
|
||||
trans: ProvidesUserContext,
|
||||
trans: SessionRequestContext,
|
||||
user: User,
|
||||
job=None,
|
||||
context: dict[str, Any] | None = None,
|
||||
|
||||
@@ -14,6 +14,7 @@ For deterministic tests without LLM, see test_static_agent_backend.py.
|
||||
import asyncio
|
||||
import logging
|
||||
import os
|
||||
from typing import cast
|
||||
|
||||
import pytest
|
||||
from fastmcp import (
|
||||
@@ -26,6 +27,7 @@ from galaxy.agents.operations import AgentOperationsManager
|
||||
from galaxy.managers.context import ProvidesUserContext
|
||||
from galaxy.util.unittest_utils import pytestmark_live_llm
|
||||
from galaxy.webapps.galaxy.api.mcp import get_mcp_app
|
||||
from galaxy.work.context import SessionRequestContext
|
||||
from galaxy_test.base.populators import (
|
||||
DatasetPopulator,
|
||||
TOOL_WITH_SHELL_COMMAND,
|
||||
@@ -216,7 +218,8 @@ class TestAgentOperationsManagerEncoding(AgentIntegrationTestCase):
|
||||
def user_is_admin(self):
|
||||
return False
|
||||
|
||||
trans = MinimalTrans(self._app)
|
||||
# the double only needs security.encode_id for _encode_ids_in_response
|
||||
trans = cast(SessionRequestContext, MinimalTrans(self._app))
|
||||
return AgentOperationsManager(app=self._app, trans=trans)
|
||||
|
||||
def test_encode_ids_helper_encodes_nested_ids(self):
|
||||
|
||||
Reference in New Issue
Block a user