chore: Add credential service to runner and invocation context

PiperOrigin-RevId: 772697298
This commit is contained in:
Xiang (Sean) Zhou
2025-06-17 18:08:43 -07:00
committed by Copybara-Service
parent 6d174eba30
commit 5f89a469ec
6 changed files with 33 additions and 6 deletions
@@ -22,6 +22,7 @@ from pydantic import BaseModel
from pydantic import ConfigDict
from ..artifacts.base_artifact_service import BaseArtifactService
from ..auth.credential_service.base_credential_service import BaseCredentialService
from ..memory.base_memory_service import BaseMemoryService
from ..sessions.base_session_service import BaseSessionService
from ..sessions.session import Session
@@ -115,6 +116,7 @@ class InvocationContext(BaseModel):
artifact_service: Optional[BaseArtifactService] = None
session_service: BaseSessionService
memory_service: Optional[BaseMemoryService] = None
credential_service: Optional[BaseCredentialService] = None
invocation_id: str
"""The id of this invocation context. Readonly."""
@@ -19,12 +19,12 @@ from abc import abstractmethod
from typing import Optional
from ...tools.tool_context import ToolContext
from ...utils.feature_decorator import working_in_progress
from ...utils.feature_decorator import experimental
from ..auth_credential import AuthCredential
from ..auth_tool import AuthConfig
@working_in_progress("Implementation are in progress. Don't use it for now.")
@experimental
class BaseCredentialService(ABC):
"""Abstract class for Service that loads / saves tool credentials from / to
the backend credential store."""
+10
View File
@@ -24,6 +24,8 @@ from pydantic import BaseModel
from ..agents.llm_agent import LlmAgent
from ..artifacts import BaseArtifactService
from ..artifacts import InMemoryArtifactService
from ..auth.credential_service.base_credential_service import BaseCredentialService
from ..auth.credential_service.in_memory_credential_service import InMemoryCredentialService
from ..runners import Runner
from ..sessions.base_session_service import BaseSessionService
from ..sessions.in_memory_session_service import InMemorySessionService
@@ -43,6 +45,7 @@ async def run_input_file(
root_agent: LlmAgent,
artifact_service: BaseArtifactService,
session_service: BaseSessionService,
credential_service: BaseCredentialService,
input_path: str,
) -> Session:
runner = Runner(
@@ -50,6 +53,7 @@ async def run_input_file(
agent=root_agent,
artifact_service=artifact_service,
session_service=session_service,
credential_service=credential_service,
)
with open(input_path, 'r', encoding='utf-8') as f:
input_file = InputFile.model_validate_json(f.read())
@@ -75,12 +79,14 @@ async def run_interactively(
artifact_service: BaseArtifactService,
session: Session,
session_service: BaseSessionService,
credential_service: BaseCredentialService,
) -> None:
runner = Runner(
app_name=session.app_name,
agent=root_agent,
artifact_service=artifact_service,
session_service=session_service,
credential_service=credential_service,
)
while True:
query = input('[user]: ')
@@ -125,6 +131,7 @@ async def run_cli(
artifact_service = InMemoryArtifactService()
session_service = InMemorySessionService()
credential_service = InMemoryCredentialService()
user_id = 'test_user'
session = await session_service.create_session(
@@ -141,6 +148,7 @@ async def run_cli(
root_agent=root_agent,
artifact_service=artifact_service,
session_service=session_service,
credential_service=credential_service,
input_path=input_file,
)
elif saved_session_file:
@@ -163,6 +171,7 @@ async def run_cli(
artifact_service,
session,
session_service,
credential_service,
)
else:
click.echo(f'Running agent {root_agent.name}, type exit to exit.')
@@ -171,6 +180,7 @@ async def run_cli(
artifact_service,
session,
session_service,
credential_service,
)
if save_session:
+5
View File
@@ -57,6 +57,7 @@ from ..agents.llm_agent import Agent
from ..agents.run_config import StreamingMode
from ..artifacts.gcs_artifact_service import GcsArtifactService
from ..artifacts.in_memory_artifact_service import InMemoryArtifactService
from ..auth.credential_service.in_memory_credential_service import InMemoryCredentialService
from ..errors.not_found_error import NotFoundError
from ..evaluation.eval_case import EvalCase
from ..evaluation.eval_case import SessionInput
@@ -305,6 +306,9 @@ def get_fast_api_app(
else:
artifact_service = InMemoryArtifactService()
# Build the Credential service
credential_service = InMemoryCredentialService()
# initialize Agent Loader
agent_loader = AgentLoader(agents_dir)
@@ -929,6 +933,7 @@ def get_fast_api_app(
artifact_service=artifact_service,
session_service=session_service,
memory_service=memory_service,
credential_service=credential_service,
)
runner_dict[app_name] = runner
return runner
+6 -1
View File
@@ -17,7 +17,6 @@ from __future__ import annotations
import asyncio
import logging
import queue
import threading
from typing import AsyncGenerator
from typing import Generator
from typing import Optional
@@ -34,6 +33,7 @@ from .agents.llm_agent import LlmAgent
from .agents.run_config import RunConfig
from .artifacts.base_artifact_service import BaseArtifactService
from .artifacts.in_memory_artifact_service import InMemoryArtifactService
from .auth.credential_service.base_credential_service import BaseCredentialService
from .code_executors.built_in_code_executor import BuiltInCodeExecutor
from .events.event import Event
from .memory.base_memory_service import BaseMemoryService
@@ -73,6 +73,8 @@ class Runner:
"""The session service for the runner."""
memory_service: Optional[BaseMemoryService] = None
"""The memory service for the runner."""
credential_service: Optional[BaseCredentialService] = None
"""The credential service for the runner."""
def __init__(
self,
@@ -82,6 +84,7 @@ class Runner:
artifact_service: Optional[BaseArtifactService] = None,
session_service: BaseSessionService,
memory_service: Optional[BaseMemoryService] = None,
credential_service: Optional[BaseCredentialService] = None,
):
"""Initializes the Runner.
@@ -97,6 +100,7 @@ class Runner:
self.artifact_service = artifact_service
self.session_service = session_service
self.memory_service = memory_service
self.credential_service = credential_service
def run(
self,
@@ -418,6 +422,7 @@ class Runner:
artifact_service=self.artifact_service,
session_service=self.session_service,
memory_service=self.memory_service,
credential_service=self.credential_service,
invocation_id=invocation_id,
agent=self.agent,
session=session,
+8 -3
View File
@@ -129,6 +129,7 @@ async def test_run_input_file_outputs(
artifact_service = cli.InMemoryArtifactService()
session_service = cli.InMemorySessionService()
credential_service = cli.InMemoryCredentialService()
dummy_root = types.SimpleNamespace(name="root")
session = await cli.run_input_file(
@@ -137,6 +138,7 @@ async def test_run_input_file_outputs(
root_agent=dummy_root,
artifact_service=artifact_service,
session_service=session_service,
credential_service=credential_service,
input_path=str(input_path),
)
@@ -199,9 +201,10 @@ async def test_run_interactively_whitespace_and_exit(
) -> None:
"""run_interactively should skip blank input, echo once, then exit."""
# make a session that belongs to dummy agent
svc = cli.InMemorySessionService()
sess = await svc.create_session(app_name="dummy", user_id="u")
session_service = cli.InMemorySessionService()
sess = await session_service.create_session(app_name="dummy", user_id="u")
artifact_service = cli.InMemoryArtifactService()
credential_service = cli.InMemoryCredentialService()
root_agent = types.SimpleNamespace(name="root")
# fake user input: blank -> 'hello' -> 'exit'
@@ -212,7 +215,9 @@ async def test_run_interactively_whitespace_and_exit(
echoed: list[str] = []
monkeypatch.setattr(click, "echo", lambda msg: echoed.append(msg))
await cli.run_interactively(root_agent, artifact_service, sess, svc)
await cli.run_interactively(
root_agent, artifact_service, sess, session_service, credential_service
)
# verify: assistant echoed once with 'echo:hello'
assert any("echo:hello" in m for m in echoed)