fix minor typos

This commit is contained in:
trandangtrungduc
2025-04-18 11:02:10 +07:00
101 changed files with 2123 additions and 258 deletions
+2 -2
View File
@@ -29,11 +29,11 @@ jobs:
run: | run: |
uv venv .venv uv venv .venv
source .venv/bin/activate source .venv/bin/activate
uv sync --extra test uv sync --extra test --extra eval
- name: Run unit tests with pytest - name: Run unit tests with pytest
run: | run: |
source .venv/bin/activate source .venv/bin/activate
pytest tests/unittests \ pytest tests/unittests \
--ignore=tests/unittests/tools/google_api_tool/test_googleapi_to_openapi_converter.py \ --ignore=tests/unittests/tools/google_api_tool/test_googleapi_to_openapi_converter.py \
--ignore=tests/unittests/artifacts/test_artifact_service.py --ignore=tests/unittests/artifacts/test_artifact_service.py
+1
View File
@@ -3,6 +3,7 @@
## 0.1.0 ## 0.1.0
### Features ### Features
* Initial release of the Agent Development Kit (ADK). * Initial release of the Agent Development Kit (ADK).
* Multi-agent, agent-as-workflow, and custom agent support * Multi-agent, agent-as-workflow, and custom agent support
* Tool authentication support * Tool authentication support
+8 -5
View File
@@ -5,9 +5,9 @@
[![r/agentdevelopmentkit](https://img.shields.io/badge/Reddit-r%2Fagentdevelopmentkit-FF4500?style=flat&logo=reddit&logoColor=white)](https://www.reddit.com/r/agentdevelopmentkit/) [![r/agentdevelopmentkit](https://img.shields.io/badge/Reddit-r%2Fagentdevelopmentkit-FF4500?style=flat&logo=reddit&logoColor=white)](https://www.reddit.com/r/agentdevelopmentkit/)
<html> <html>
<h1 align="center"> <h2 align="center">
<img src="assets/agent-development-kit.png" width="256"/> <img src="https://raw.githubusercontent.com/google/adk-python/main/assets/agent-development-kit.png" width="256"/>
</h1> </h2>
<h3 align="center"> <h3 align="center">
An open-source, code-first Python toolkit for building, evaluating, and deploying sophisticated AI agents with flexibility and control. An open-source, code-first Python toolkit for building, evaluating, and deploying sophisticated AI agents with flexibility and control.
</h3> </h3>
@@ -50,6 +50,7 @@ You can install the ADK using `pip`:
```bash ```bash
pip install google-adk pip install google-adk
``` ```
## 📚 Documentation ## 📚 Documentation
Explore the full documentation for detailed guides on building, evaluating, and Explore the full documentation for detailed guides on building, evaluating, and
@@ -60,6 +61,7 @@ deploying agents:
## 🏁 Feature Highlight ## 🏁 Feature Highlight
### Define a single agent: ### Define a single agent:
```python ```python
from google.adk.agents import Agent from google.adk.agents import Agent
from google.adk.tools import google_search from google.adk.tools import google_search
@@ -74,7 +76,9 @@ root_agent = Agent(
``` ```
### Define a multi-agent system: ### Define a multi-agent system:
Define a multi-agent system with coordinator agent, greeter agent, and task execution agent. Then ADK engine and the model will guide the agents works together to accomplish the task. Define a multi-agent system with coordinator agent, greeter agent, and task execution agent. Then ADK engine and the model will guide the agents works together to accomplish the task.
```python ```python
from google.adk.agents import LlmAgent, BaseAgent from google.adk.agents import LlmAgent, BaseAgent
@@ -92,14 +96,13 @@ coordinator = LlmAgent(
task_executor task_executor
] ]
) )
``` ```
### Development UI ### Development UI
A built-in development UI to help you test, evaluate, debug, and showcase your agent(s). A built-in development UI to help you test, evaluate, debug, and showcase your agent(s).
<img src="assets/adk-web-dev-ui-function-call.png"/> <img src="https://raw.githubusercontent.com/google/adk-python/main/assets/adk-web-dev-ui-function-call.png"/>
### Evaluate Agents ### Evaluate Agents
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
+2
View File
@@ -82,6 +82,8 @@ async def run_interactively(
) )
while True: while True:
query = input('user: ') query = input('user: ')
if not query or not query.strip():
continue
if query == 'exit': if query == 'exit':
break break
async for event in runner.run_async( async for event in runner.run_async(
+279
View File
@@ -0,0 +1,279 @@
# Copyright 2025 Google LLC
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
import os
import subprocess
from typing import Optional
from typing import Tuple
import click
_INIT_PY_TEMPLATE = """\
from . import agent
"""
_AGENT_PY_TEMPLATE = """\
from google.adk.agents import Agent
root_agent = Agent(
model='{model_name}',
name='root_agent',
description='A helpful assistant for user questions.',
instruction='Answer user questions to the best of your knowledge',
)
"""
_GOOGLE_API_MSG = """
Don't have API Key? Create one in AI Studio: https://aistudio.google.com/apikey
"""
_GOOGLE_CLOUD_SETUP_MSG = """
You need an existing Google Cloud account and project, check out this link for details:
https://google.github.io/adk-docs/get-started/quickstart/#gemini---google-cloud-vertex-ai
"""
_OTHER_MODEL_MSG = """
Please see below guide to configure other models:
https://google.github.io/adk-docs/agents/models
"""
_SUCCESS_MSG = """
Agent created in {agent_folder}:
- .env
- __init__.py
- agent.py
"""
def _get_gcp_project_from_gcloud() -> str:
"""Uses gcloud to get default project."""
try:
result = subprocess.run(
["gcloud", "config", "get-value", "project"],
capture_output=True,
text=True,
check=True,
)
return result.stdout.strip()
except (subprocess.CalledProcessError, FileNotFoundError):
return ""
def _get_gcp_region_from_gcloud() -> str:
"""Uses gcloud to get default region."""
try:
result = subprocess.run(
["gcloud", "config", "get-value", "compute/region"],
capture_output=True,
text=True,
check=True,
)
return result.stdout.strip()
except (subprocess.CalledProcessError, FileNotFoundError):
return ""
def _prompt_str(
prompt_prefix: str,
*,
prior_msg: Optional[str] = None,
default_value: Optional[str] = None,
) -> str:
if prior_msg:
click.secho(prior_msg, fg="green")
while True:
value: str = click.prompt(
prompt_prefix, default=default_value or None, type=str
)
if value and value.strip():
return value.strip()
def _prompt_for_google_cloud(
google_cloud_project: Optional[str],
) -> str:
"""Prompts user for Google Cloud project ID."""
google_cloud_project = (
google_cloud_project
or os.environ.get("GOOGLE_CLOUD_PROJECT", None)
or _get_gcp_project_from_gcloud()
)
google_cloud_project = _prompt_str(
"Enter Google Cloud project ID", default_value=google_cloud_project
)
return google_cloud_project
def _prompt_for_google_cloud_region(
google_cloud_region: Optional[str],
) -> str:
"""Prompts user for Google Cloud region."""
google_cloud_region = (
google_cloud_region
or os.environ.get("GOOGLE_CLOUD_LOCATION", None)
or _get_gcp_region_from_gcloud()
)
google_cloud_region = _prompt_str(
"Enter Google Cloud region",
default_value=google_cloud_region or "us-central1",
)
return google_cloud_region
def _prompt_for_google_api_key(
google_api_key: Optional[str],
) -> str:
"""Prompts user for Google API key."""
google_api_key = google_api_key or os.environ.get("GOOGLE_API_KEY", None)
google_api_key = _prompt_str(
"Enter Google API key",
prior_msg=_GOOGLE_API_MSG,
default_value=google_api_key,
)
return google_api_key
def _generate_files(
agent_folder: str,
*,
google_api_key: Optional[str] = None,
google_cloud_project: Optional[str] = None,
google_cloud_region: Optional[str] = None,
model: Optional[str] = None,
):
"""Generates a folder name for the agent."""
os.makedirs(agent_folder, exist_ok=True)
dotenv_file_path = os.path.join(agent_folder, ".env")
init_file_path = os.path.join(agent_folder, "__init__.py")
agent_file_path = os.path.join(agent_folder, "agent.py")
with open(dotenv_file_path, "w", encoding="utf-8") as f:
lines = []
if google_api_key:
lines.append("GOOGLE_GENAI_USE_VERTEXAI=0")
elif google_cloud_project and google_cloud_region:
lines.append("GOOGLE_GENAI_USE_VERTEXAI=1")
if google_api_key:
lines.append(f"GOOGLE_API_KEY={google_api_key}")
if google_cloud_project:
lines.append(f"GOOGLE_CLOUD_PROJECT={google_cloud_project}")
if google_cloud_region:
lines.append(f"GOOGLE_CLOUD_LOCATION={google_cloud_region}")
f.write("\n".join(lines))
with open(init_file_path, "w", encoding="utf-8") as f:
f.write(_INIT_PY_TEMPLATE)
with open(agent_file_path, "w", encoding="utf-8") as f:
f.write(_AGENT_PY_TEMPLATE.format(model_name=model))
click.secho(
_SUCCESS_MSG.format(agent_folder=agent_folder),
fg="green",
)
def _prompt_for_model() -> str:
model_choice = click.prompt(
"""\
Choose a model for the root agent:
1. gemini-2.0-flash-001
2. Other models (fill later)
Choose model""",
type=click.Choice(["1", "2"]),
)
if model_choice == "1":
return "gemini-2.0-flash-001"
else:
click.secho(_OTHER_MODEL_MSG, fg="green")
return "<FILL_IN_MODEL>"
def _prompt_to_choose_backend(
google_api_key: Optional[str],
google_cloud_project: Optional[str],
google_cloud_region: Optional[str],
) -> Tuple[Optional[str], Optional[str], Optional[str]]:
"""Prompts user to choose backend.
Returns:
A tuple of (google_api_key, google_cloud_project, google_cloud_region).
"""
backend_choice = click.prompt(
"1. Google AI\n2. Vertex AI\nChoose a backend",
type=click.Choice(["1", "2"]),
)
if backend_choice == "1":
google_api_key = _prompt_for_google_api_key(google_api_key)
elif backend_choice == "2":
click.secho(_GOOGLE_CLOUD_SETUP_MSG, fg="green")
google_cloud_project = _prompt_for_google_cloud(google_cloud_project)
google_cloud_region = _prompt_for_google_cloud_region(google_cloud_region)
return google_api_key, google_cloud_project, google_cloud_region
def run_cmd(
agent_name: str,
*,
model: Optional[str],
google_api_key: Optional[str],
google_cloud_project: Optional[str],
google_cloud_region: Optional[str],
):
"""Runs `adk create` command to create agent template.
Args:
agent_name: str, The name of the agent.
google_api_key: Optional[str], The Google API key for using Google AI as
backend.
google_cloud_project: Optional[str], The Google Cloud project for using
VertexAI as backend.
google_cloud_region: Optional[str], The Google Cloud region for using
VertexAI as backend.
"""
agent_folder = os.path.join(os.getcwd(), agent_name)
# check folder doesn't exist or it's empty. Otherwise, throw
if os.path.exists(agent_folder) and os.listdir(agent_folder):
# Prompt user whether to override existing files using click
if not click.confirm(
f"Non-empty folder already exist: '{agent_folder}'\n"
"Override existing content?",
default=False,
):
raise click.Abort()
if not model:
model = _prompt_for_model()
if not google_api_key and not (google_cloud_project and google_cloud_region):
if model.startswith("gemini"):
google_api_key, google_cloud_project, google_cloud_region = (
_prompt_to_choose_backend(
google_api_key, google_cloud_project, google_cloud_region
)
)
_generate_files(
agent_folder,
google_api_key=google_api_key,
google_cloud_project=google_cloud_project,
google_cloud_region=google_cloud_region,
model=model,
)
+7 -3
View File
@@ -82,8 +82,9 @@ def to_cloud_run(
app_name: str, app_name: str,
temp_folder: str, temp_folder: str,
port: int, port: int,
with_cloud_trace: bool, trace_to_cloud: bool,
with_ui: bool, with_ui: bool,
verbosity: str,
): ):
"""Deploys an agent to Google Cloud Run. """Deploys an agent to Google Cloud Run.
@@ -108,8 +109,9 @@ def to_cloud_run(
app_name: The name of the app, by default, it's basename of `agent_folder`. app_name: The name of the app, by default, it's basename of `agent_folder`.
temp_folder: The temp folder for the generated Cloud Run source files. temp_folder: The temp folder for the generated Cloud Run source files.
port: The port of the ADK api server. port: The port of the ADK api server.
with_cloud_trace: Whether to enable Cloud Trace. trace_to_cloud: Whether to enable Cloud Trace.
with_ui: Whether to deploy with UI. with_ui: Whether to deploy with UI.
verbosity: The verbosity level of the CLI.
""" """
app_name = app_name or os.path.basename(agent_folder) app_name = app_name or os.path.basename(agent_folder)
@@ -142,7 +144,7 @@ def to_cloud_run(
port=port, port=port,
command='web' if with_ui else 'api_server', command='web' if with_ui else 'api_server',
install_agent_deps=install_agent_deps, install_agent_deps=install_agent_deps,
trace_to_cloud_option='--trace_to_cloud' if with_cloud_trace else '', trace_to_cloud_option='--trace_to_cloud' if trace_to_cloud else '',
) )
dockerfile_path = os.path.join(temp_folder, 'Dockerfile') dockerfile_path = os.path.join(temp_folder, 'Dockerfile')
os.makedirs(temp_folder, exist_ok=True) os.makedirs(temp_folder, exist_ok=True)
@@ -169,6 +171,8 @@ def to_cloud_run(
*region_options, *region_options,
'--port', '--port',
str(port), str(port),
'--verbosity',
verbosity,
'--labels', '--labels',
'created-by=adk', 'created-by=adk',
], ],
+71 -11
View File
@@ -24,6 +24,7 @@ import click
from fastapi import FastAPI from fastapi import FastAPI
import uvicorn import uvicorn
from . import cli_create
from . import cli_deploy from . import cli_deploy
from .cli import run_cli from .cli import run_cli
from .cli_eval import MISSING_EVAL_DEPENDENCIES_MESSAGE from .cli_eval import MISSING_EVAL_DEPENDENCIES_MESSAGE
@@ -42,10 +43,59 @@ def main():
@main.group() @main.group()
def deploy(): def deploy():
"""Deploy Agent.""" """Deploys agent to hosted environments."""
pass pass
@main.command("create")
@click.option(
"--model",
type=str,
help="Optional. The model used for the root agent.",
)
@click.option(
"--api_key",
type=str,
help=(
"Optional. The API Key needed to access the model, e.g. Google AI API"
" Key."
),
)
@click.option(
"--project",
type=str,
help="Optional. The Google Cloud Project for using VertexAI as backend.",
)
@click.option(
"--region",
type=str,
help="Optional. The Google Cloud Region for using VertexAI as backend.",
)
@click.argument("app_name", type=str, required=True)
def cli_create_cmd(
app_name: str,
model: Optional[str],
api_key: Optional[str],
project: Optional[str],
region: Optional[str],
):
"""Creates a new app in the current folder with prepopulated agent template.
APP_NAME: required, the folder of the agent source code.
Example:
adk create path/to/my_app
"""
cli_create.run_cmd(
app_name,
model=model,
google_api_key=api_key,
google_cloud_project=project,
google_cloud_region=region,
)
@main.command("run") @main.command("run")
@click.option( @click.option(
"--save_session", "--save_session",
@@ -62,7 +112,7 @@ def deploy():
), ),
) )
def cli_run(agent: str, save_session: bool): def cli_run(agent: str, save_session: bool):
"""Run an interactive CLI for a certain agent. """Runs an interactive CLI for a certain agent.
AGENT: The path to the agent source code folder. AGENT: The path to the agent source code folder.
@@ -150,7 +200,7 @@ def cli_eval(
EvalMetric(metric_name=metric_name, threshold=threshold) EvalMetric(metric_name=metric_name, threshold=threshold)
) )
print(f"Using evaluation criteria: {evaluation_criteria}") print(f"Using evaluation creiteria: {evaluation_criteria}")
root_agent = get_root_agent(agent_module_file_path) root_agent = get_root_agent(agent_module_file_path)
reset_func = try_get_reset_func(agent_module_file_path) reset_func = try_get_reset_func(agent_module_file_path)
@@ -244,7 +294,7 @@ def cli_eval(
type=click.Path( type=click.Path(
exists=True, dir_okay=True, file_okay=False, resolve_path=True exists=True, dir_okay=True, file_okay=False, resolve_path=True
), ),
default=os.getcwd(), default=os.getcwd,
) )
def cli_web( def cli_web(
agents_dir: str, agents_dir: str,
@@ -255,7 +305,7 @@ def cli_web(
port: int = 8000, port: int = 8000,
trace_to_cloud: bool = False, trace_to_cloud: bool = False,
): ):
"""Start a FastAPI server with Web UI for agents. """Starts a FastAPI server with Web UI for agents.
AGENTS_DIR: The directory of agents, where each sub-directory is a single AGENTS_DIR: The directory of agents, where each sub-directory is a single
agent, containing at least `__init__.py` and `agent.py` files. agent, containing at least `__init__.py` and `agent.py` files.
@@ -274,7 +324,7 @@ def cli_web(
@asynccontextmanager @asynccontextmanager
async def _lifespan(app: FastAPI): async def _lifespan(app: FastAPI):
click.secho( click.secho(
f"""\ f"""
+-----------------------------------------------------------------------------+ +-----------------------------------------------------------------------------+
| ADK Web Server started | | ADK Web Server started |
| | | |
@@ -285,7 +335,7 @@ def cli_web(
) )
yield # Startup is done, now app is running yield # Startup is done, now app is running
click.secho( click.secho(
"""\ """
+-----------------------------------------------------------------------------+ +-----------------------------------------------------------------------------+
| ADK Web Server shutting down... | | ADK Web Server shutting down... |
+-----------------------------------------------------------------------------+ +-----------------------------------------------------------------------------+
@@ -378,7 +428,7 @@ def cli_api_server(
port: int = 8000, port: int = 8000,
trace_to_cloud: bool = False, trace_to_cloud: bool = False,
): ):
"""Start a FastAPI server for agents. """Starts a FastAPI server for agents.
AGENTS_DIR: The directory of agents, where each sub-directory is a single AGENTS_DIR: The directory of agents, where each sub-directory is a single
agent, containing at least `__init__.py` and `agent.py` files. agent, containing at least `__init__.py` and `agent.py` files.
@@ -452,7 +502,7 @@ def cli_api_server(
help="Optional. The port of the ADK API server (default: 8000).", help="Optional. The port of the ADK API server (default: 8000).",
) )
@click.option( @click.option(
"--with_cloud_trace", "--trace_to_cloud",
type=bool, type=bool,
is_flag=True, is_flag=True,
show_default=True, show_default=True,
@@ -483,6 +533,14 @@ def cli_api_server(
" (default: a timestamped folder in the system temp directory)." " (default: a timestamped folder in the system temp directory)."
), ),
) )
@click.option(
"--verbosity",
type=click.Choice(
["debug", "info", "warning", "error", "critical"], case_sensitive=False
),
default="WARNING",
help="Optional. Override the default verbosity level.",
)
@click.argument( @click.argument(
"agent", "agent",
type=click.Path( type=click.Path(
@@ -497,8 +555,9 @@ def cli_deploy_cloud_run(
app_name: str, app_name: str,
temp_folder: str, temp_folder: str,
port: int, port: int,
with_cloud_trace: bool, trace_to_cloud: bool,
with_ui: bool, with_ui: bool,
verbosity: str,
): ):
"""Deploys an agent to Cloud Run. """Deploys an agent to Cloud Run.
@@ -517,8 +576,9 @@ def cli_deploy_cloud_run(
app_name=app_name, app_name=app_name,
temp_folder=temp_folder, temp_folder=temp_folder,
port=port, port=port,
with_cloud_trace=with_cloud_trace, trace_to_cloud=trace_to_cloud,
with_ui=with_ui, with_ui=with_ui,
verbosity=verbosity,
) )
except Exception as e: except Exception as e:
click.secho(f"Deploy failed: {e}", fg="red", err=True) click.secho(f"Deploy failed: {e}", fg="red", err=True)
+52 -20
View File
@@ -13,7 +13,9 @@
# limitations under the License. # limitations under the License.
import asyncio import asyncio
from contextlib import asynccontextmanager
import importlib import importlib
import inspect
import json import json
import logging import logging
import os import os
@@ -28,6 +30,7 @@ from typing import Literal
from typing import Optional from typing import Optional
import click import click
from click import Tuple
from fastapi import FastAPI from fastapi import FastAPI
from fastapi import HTTPException from fastapi import HTTPException
from fastapi import Query from fastapi import Query
@@ -56,6 +59,7 @@ from ..agents.llm_agent import Agent
from ..agents.run_config import StreamingMode from ..agents.run_config import StreamingMode
from ..artifacts import InMemoryArtifactService from ..artifacts import InMemoryArtifactService
from ..events.event import Event from ..events.event import Event
from ..memory.in_memory_memory_service import InMemoryMemoryService
from ..runners import Runner from ..runners import Runner
from ..sessions.database_session_service import DatabaseSessionService from ..sessions.database_session_service import DatabaseSessionService
from ..sessions.in_memory_session_service import InMemorySessionService from ..sessions.in_memory_session_service import InMemorySessionService
@@ -143,11 +147,8 @@ def get_fast_api_app(
provider.add_span_processor( provider.add_span_processor(
export.SimpleSpanProcessor(ApiServerSpanExporter(trace_dict)) export.SimpleSpanProcessor(ApiServerSpanExporter(trace_dict))
) )
envs.load_dotenv() if trace_to_cloud:
enable_cloud_tracing = trace_to_cloud or os.environ.get( envs.load_dotenv_for_agent("", agent_dir)
"ADK_TRACE_TO_CLOUD", "0"
).lower() in ["1", "true"]
if enable_cloud_tracing:
if project_id := os.environ.get("GOOGLE_CLOUD_PROJECT", None): if project_id := os.environ.get("GOOGLE_CLOUD_PROJECT", None):
processor = export.BatchSpanProcessor( processor = export.BatchSpanProcessor(
CloudTraceSpanExporter(project_id=project_id) CloudTraceSpanExporter(project_id=project_id)
@@ -161,8 +162,22 @@ def get_fast_api_app(
trace.set_tracer_provider(provider) trace.set_tracer_provider(provider)
exit_stacks = []
@asynccontextmanager
async def internal_lifespan(app: FastAPI):
if lifespan:
async with lifespan(app) as lifespan_context:
yield
if exit_stacks:
for stack in exit_stacks:
await stack.aclose()
else:
yield
# Run the FastAPI server. # Run the FastAPI server.
app = FastAPI(lifespan=lifespan) app = FastAPI(lifespan=internal_lifespan)
if allow_origins: if allow_origins:
app.add_middleware( app.add_middleware(
@@ -181,6 +196,7 @@ def get_fast_api_app(
# Build the Artifact service # Build the Artifact service
artifact_service = InMemoryArtifactService() artifact_service = InMemoryArtifactService()
memory_service = InMemoryMemoryService()
# Build the Session service # Build the Session service
agent_engine_id = "" agent_engine_id = ""
@@ -358,7 +374,7 @@ def get_fast_api_app(
"/apps/{app_name}/eval_sets/{eval_set_id}/add_session", "/apps/{app_name}/eval_sets/{eval_set_id}/add_session",
response_model_exclude_none=True, response_model_exclude_none=True,
) )
def add_session_to_eval_set( async def add_session_to_eval_set(
app_name: str, eval_set_id: str, req: AddSessionToEvalSetRequest app_name: str, eval_set_id: str, req: AddSessionToEvalSetRequest
): ):
pattern = r"^[a-zA-Z0-9_]+$" pattern = r"^[a-zA-Z0-9_]+$"
@@ -393,7 +409,9 @@ def get_fast_api_app(
test_data = evals.convert_session_to_eval_format(session) test_data = evals.convert_session_to_eval_format(session)
# Populate the session with initial session state. # Populate the session with initial session state.
initial_session_state = create_empty_state(_get_root_agent(app_name)) initial_session_state = create_empty_state(
await _get_root_agent_async(app_name)
)
eval_set_data.append({ eval_set_data.append({
"name": req.eval_id, "name": req.eval_id,
@@ -430,7 +448,7 @@ def get_fast_api_app(
"/apps/{app_name}/eval_sets/{eval_set_id}/run_eval", "/apps/{app_name}/eval_sets/{eval_set_id}/run_eval",
response_model_exclude_none=True, response_model_exclude_none=True,
) )
def run_eval( async def run_eval(
app_name: str, eval_set_id: str, req: RunEvalRequest app_name: str, eval_set_id: str, req: RunEvalRequest
) -> list[RunEvalResult]: ) -> list[RunEvalResult]:
from .cli_eval import run_evals from .cli_eval import run_evals
@@ -447,7 +465,7 @@ def get_fast_api_app(
logger.info( logger.info(
"Eval ids to run list is empty. We will all evals in the eval set." "Eval ids to run list is empty. We will all evals in the eval set."
) )
root_agent = _get_root_agent(app_name) root_agent = await _get_root_agent_async(app_name)
eval_results = list( eval_results = list(
run_evals( run_evals(
eval_set_to_evals, eval_set_to_evals,
@@ -577,7 +595,7 @@ def get_fast_api_app(
) )
if not session: if not session:
raise HTTPException(status_code=404, detail="Session not found") raise HTTPException(status_code=404, detail="Session not found")
runner = _get_runner(req.app_name) runner = await _get_runner_async(req.app_name)
events = [ events = [
event event
async for event in runner.run_async( async for event in runner.run_async(
@@ -604,7 +622,7 @@ def get_fast_api_app(
async def event_generator(): async def event_generator():
try: try:
stream_mode = StreamingMode.SSE if req.streaming else StreamingMode.NONE stream_mode = StreamingMode.SSE if req.streaming else StreamingMode.NONE
runner = _get_runner(req.app_name) runner = await _get_runner_async(req.app_name)
async for event in runner.run_async( async for event in runner.run_async(
user_id=req.user_id, user_id=req.user_id,
session_id=req.session_id, session_id=req.session_id,
@@ -630,7 +648,7 @@ def get_fast_api_app(
"/apps/{app_name}/users/{user_id}/sessions/{session_id}/events/{event_id}/graph", "/apps/{app_name}/users/{user_id}/sessions/{session_id}/events/{event_id}/graph",
response_model_exclude_none=True, response_model_exclude_none=True,
) )
def get_event_graph( async def get_event_graph(
app_name: str, user_id: str, session_id: str, event_id: str app_name: str, user_id: str, session_id: str, event_id: str
): ):
# Connect to managed session if agent_engine_id is set. # Connect to managed session if agent_engine_id is set.
@@ -647,7 +665,7 @@ def get_fast_api_app(
function_calls = event.get_function_calls() function_calls = event.get_function_calls()
function_responses = event.get_function_responses() function_responses = event.get_function_responses()
root_agent = _get_root_agent(app_name) root_agent = await _get_root_agent_async(app_name)
dot_graph = None dot_graph = None
if function_calls: if function_calls:
function_call_highlights = [] function_call_highlights = []
@@ -704,7 +722,7 @@ def get_fast_api_app(
live_request_queue = LiveRequestQueue() live_request_queue = LiveRequestQueue()
async def forward_events(): async def forward_events():
runner = _get_runner(app_name) runner = await _get_runner_async(app_name)
async for event in runner.run_live( async for event in runner.run_live(
session=session, live_request_queue=live_request_queue session=session, live_request_queue=live_request_queue
): ):
@@ -742,26 +760,40 @@ def get_fast_api_app(
for task in pending: for task in pending:
task.cancel() task.cancel()
def _get_root_agent(app_name: str) -> Agent: async def _get_root_agent_async(app_name: str) -> Agent:
"""Returns the root agent for the given app.""" """Returns the root agent for the given app."""
if app_name in root_agent_dict: if app_name in root_agent_dict:
return root_agent_dict[app_name] return root_agent_dict[app_name]
envs.load_dotenv_for_agent(os.path.basename(app_name), agent_dir)
agent_module = importlib.import_module(app_name) agent_module = importlib.import_module(app_name)
root_agent: Agent = agent_module.agent.root_agent if getattr(agent_module.agent, "root_agent"):
root_agent = agent_module.agent.root_agent
else:
raise ValueError(f'Unable to find "root_agent" from {app_name}.')
# Handle an awaitable root agent and await for the actual agent.
if inspect.isawaitable(root_agent):
try:
agent, exit_stack = await root_agent
exit_stacks.append(exit_stack)
root_agent = agent
except Exception as e:
raise RuntimeError(f"error getting root agent, {e}") from e
root_agent_dict[app_name] = root_agent root_agent_dict[app_name] = root_agent
return root_agent return root_agent
def _get_runner(app_name: str) -> Runner: async def _get_runner_async(app_name: str) -> Runner:
"""Returns the runner for the given app.""" """Returns the runner for the given app."""
envs.load_dotenv_for_agent(os.path.basename(app_name), agent_dir)
if app_name in runner_dict: if app_name in runner_dict:
return runner_dict[app_name] return runner_dict[app_name]
root_agent = _get_root_agent(app_name) root_agent = await _get_root_agent_async(app_name)
runner = Runner( runner = Runner(
app_name=agent_engine_id if agent_engine_id else app_name, app_name=agent_engine_id if agent_engine_id else app_name,
agent=root_agent, agent=root_agent,
artifact_service=artifact_service, artifact_service=artifact_service,
session_service=session_service, session_service=session_service,
memory_service=memory_service,
) )
runner_dict[app_name] = runner runner_dict[app_name] = runner
return runner return runner
-3
View File
@@ -50,8 +50,5 @@ def load_dotenv_for_agent(
agent_name, agent_name,
dotenv_file_path, dotenv_file_path,
) )
logger.info(
'Reloaded %s file for %s at %s', filename, agent_name, dotenv_file_path
)
else: else:
logger.info('No %s file found for %s', filename, agent_name) logger.info('No %s file found for %s', filename, agent_name)
@@ -106,9 +106,11 @@ class ResponseEvaluator:
eval_dataset = pd.DataFrame(flattened_queries).rename( eval_dataset = pd.DataFrame(flattened_queries).rename(
columns={"query": "prompt", "expected_tool_use": "reference_trajectory"} columns={"query": "prompt", "expected_tool_use": "reference_trajectory"}
) )
eval_task = EvalTask(dataset=eval_dataset, metrics=metrics)
eval_result = eval_task.evaluate() eval_result = ResponseEvaluator._perform_eval(
dataset=eval_dataset, metrics=metrics
)
if print_detailed_results: if print_detailed_results:
ResponseEvaluator._print_results(eval_result) ResponseEvaluator._print_results(eval_result)
return eval_result.summary_metrics return eval_result.summary_metrics
@@ -129,6 +131,16 @@ class ResponseEvaluator:
metrics.append("rouge_1") metrics.append("rouge_1")
return metrics return metrics
@staticmethod
def _perform_eval(dataset, metrics):
"""This method hides away the call to external service.
Primarily helps with unit testing.
"""
eval_task = EvalTask(dataset=dataset, metrics=metrics)
return eval_task.evaluate()
@staticmethod @staticmethod
def _print_results(eval_result): def _print_results(eval_result):
print("Evaluation Summary Metrics:", eval_result.summary_metrics) print("Evaluation Summary Metrics:", eval_result.summary_metrics)
+10 -4
View File
@@ -87,15 +87,21 @@ class _NlPlanningResponse(BaseLlmResponseProcessor):
return return
# Postprocess the LLM response. # Postprocess the LLM response.
callback_context = CallbackContext(invocation_context)
processed_parts = planner.process_planning_response( processed_parts = planner.process_planning_response(
CallbackContext(invocation_context), llm_response.content.parts callback_context, llm_response.content.parts
) )
if processed_parts: if processed_parts:
llm_response.content.parts = processed_parts llm_response.content.parts = processed_parts
# Maintain async generator behavior if callback_context.state.has_delta():
if False: # Ensures it behaves as a generator state_update_event = Event(
yield # This is a no-op but maintains generator structure invocation_id=invocation_context.invocation_id,
author=invocation_context.agent.name,
branch=invocation_context.branch,
actions=callback_context._event_actions,
)
yield state_update_event
response_processor = _NlPlanningResponse() response_processor = _NlPlanningResponse()
@@ -12,6 +12,7 @@
# See the License for the specific language governing permissions and # See the License for the specific language governing permissions and
# limitations under the License. # limitations under the License.
import base64
import copy import copy
from datetime import datetime from datetime import datetime
import json import json
@@ -20,17 +21,17 @@ from typing import Any
from typing import Optional from typing import Optional
import uuid import uuid
from sqlalchemy import Boolean
from sqlalchemy import delete from sqlalchemy import delete
from sqlalchemy import Dialect from sqlalchemy import Dialect
from sqlalchemy import ForeignKeyConstraint from sqlalchemy import ForeignKeyConstraint
from sqlalchemy import func from sqlalchemy import func
from sqlalchemy import select
from sqlalchemy import Text from sqlalchemy import Text
from sqlalchemy.dialects import postgresql from sqlalchemy.dialects import postgresql
from sqlalchemy.engine import create_engine from sqlalchemy.engine import create_engine
from sqlalchemy.engine import Engine from sqlalchemy.engine import Engine
from sqlalchemy.ext.mutable import MutableDict
from sqlalchemy.exc import ArgumentError from sqlalchemy.exc import ArgumentError
from sqlalchemy.ext.mutable import MutableDict
from sqlalchemy.inspection import inspect from sqlalchemy.inspection import inspect
from sqlalchemy.orm import DeclarativeBase from sqlalchemy.orm import DeclarativeBase
from sqlalchemy.orm import Mapped from sqlalchemy.orm import Mapped
@@ -54,6 +55,7 @@ from .base_session_service import ListSessionsResponse
from .session import Session from .session import Session
from .state import State from .state import State
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -103,7 +105,7 @@ class StorageSession(Base):
String, primary_key=True, default=lambda: str(uuid.uuid4()) String, primary_key=True, default=lambda: str(uuid.uuid4())
) )
state: Mapped[dict] = mapped_column( state: Mapped[MutableDict[str, Any]] = mapped_column(
MutableDict.as_mutable(DynamicJSON), default={} MutableDict.as_mutable(DynamicJSON), default={}
) )
@@ -134,8 +136,20 @@ class StorageEvent(Base):
author: Mapped[str] = mapped_column(String) author: Mapped[str] = mapped_column(String)
branch: Mapped[str] = mapped_column(String, nullable=True) branch: Mapped[str] = mapped_column(String, nullable=True)
timestamp: Mapped[DateTime] = mapped_column(DateTime(), default=func.now()) timestamp: Mapped[DateTime] = mapped_column(DateTime(), default=func.now())
content: Mapped[dict] = mapped_column(DynamicJSON) content: Mapped[dict[str, Any]] = mapped_column(DynamicJSON)
actions: Mapped[dict] = mapped_column(PickleType) actions: Mapped[MutableDict[str, Any]] = mapped_column(PickleType)
long_running_tool_ids_json: Mapped[Optional[str]] = mapped_column(
Text, nullable=True
)
grounding_metadata: Mapped[dict[str, Any]] = mapped_column(
DynamicJSON, nullable=True
)
partial: Mapped[bool] = mapped_column(Boolean, nullable=True)
turn_complete: Mapped[bool] = mapped_column(Boolean, nullable=True)
error_code: Mapped[str] = mapped_column(String, nullable=True)
error_message: Mapped[str] = mapped_column(String, nullable=True)
interrupted: Mapped[bool] = mapped_column(Boolean, nullable=True)
storage_session: Mapped[StorageSession] = relationship( storage_session: Mapped[StorageSession] = relationship(
"StorageSession", "StorageSession",
@@ -150,13 +164,28 @@ class StorageEvent(Base):
), ),
) )
@property
def long_running_tool_ids(self) -> set[str]:
return (
set(json.loads(self.long_running_tool_ids_json))
if self.long_running_tool_ids_json
else set()
)
@long_running_tool_ids.setter
def long_running_tool_ids(self, value: set[str]):
if value is None:
self.long_running_tool_ids_json = None
else:
self.long_running_tool_ids_json = json.dumps(list(value))
class StorageAppState(Base): class StorageAppState(Base):
"""Represents an app state stored in the database.""" """Represents an app state stored in the database."""
__tablename__ = "app_states" __tablename__ = "app_states"
app_name: Mapped[str] = mapped_column(String, primary_key=True) app_name: Mapped[str] = mapped_column(String, primary_key=True)
state: Mapped[dict] = mapped_column( state: Mapped[MutableDict[str, Any]] = mapped_column(
MutableDict.as_mutable(DynamicJSON), default={} MutableDict.as_mutable(DynamicJSON), default={}
) )
update_time: Mapped[DateTime] = mapped_column( update_time: Mapped[DateTime] = mapped_column(
@@ -170,7 +199,7 @@ class StorageUserState(Base):
app_name: Mapped[str] = mapped_column(String, primary_key=True) app_name: Mapped[str] = mapped_column(String, primary_key=True)
user_id: Mapped[str] = mapped_column(String, primary_key=True) user_id: Mapped[str] = mapped_column(String, primary_key=True)
state: Mapped[dict] = mapped_column( state: Mapped[MutableDict[str, Any]] = mapped_column(
MutableDict.as_mutable(DynamicJSON), default={} MutableDict.as_mutable(DynamicJSON), default={}
) )
update_time: Mapped[DateTime] = mapped_column( update_time: Mapped[DateTime] = mapped_column(
@@ -295,7 +324,6 @@ class DatabaseSessionService(BaseSessionService):
last_update_time=storage_session.update_time.timestamp(), last_update_time=storage_session.update_time.timestamp(),
) )
return session return session
return None
@override @override
def get_session( def get_session(
@@ -309,7 +337,6 @@ class DatabaseSessionService(BaseSessionService):
# 1. Get the storage session entry from session table # 1. Get the storage session entry from session table
# 2. Get all the events based on session id and filtering config # 2. Get all the events based on session id and filtering config
# 3. Convert and return the session # 3. Convert and return the session
session: Session = None
with self.DatabaseSessionFactory() as sessionFactory: with self.DatabaseSessionFactory() as sessionFactory:
storage_session = sessionFactory.get( storage_session = sessionFactory.get(
StorageSession, (app_name, user_id, session_id) StorageSession, (app_name, user_id, session_id)
@@ -356,13 +383,19 @@ class DatabaseSessionService(BaseSessionService):
author=e.author, author=e.author,
branch=e.branch, branch=e.branch,
invocation_id=e.invocation_id, invocation_id=e.invocation_id,
content=e.content, content=_decode_content(e.content),
actions=e.actions, actions=e.actions,
timestamp=e.timestamp.timestamp(), timestamp=e.timestamp.timestamp(),
long_running_tool_ids=e.long_running_tool_ids,
grounding_metadata=e.grounding_metadata,
partial=e.partial,
turn_complete=e.turn_complete,
error_code=e.error_code,
error_message=e.error_message,
interrupted=e.interrupted,
) )
for e in storage_events for e in storage_events
] ]
return session return session
@override @override
@@ -387,7 +420,6 @@ class DatabaseSessionService(BaseSessionService):
) )
sessions.append(session) sessions.append(session)
return ListSessionsResponse(sessions=sessions) return ListSessionsResponse(sessions=sessions)
raise ValueError("Failed to retrieve sessions.")
@override @override
def delete_session( def delete_session(
@@ -406,7 +438,7 @@ class DatabaseSessionService(BaseSessionService):
def append_event(self, session: Session, event: Event) -> Event: def append_event(self, session: Session, event: Event) -> Event:
logger.info(f"Append event: {event} to session {session.id}") logger.info(f"Append event: {event} to session {session.id}")
if event.partial and not event.content: if event.partial:
return event return event
# 1. Check if timestamp is stale # 1. Check if timestamp is stale
@@ -455,19 +487,34 @@ class DatabaseSessionService(BaseSessionService):
storage_user_state.state = user_state storage_user_state.state = user_state
storage_session.state = session_state storage_session.state = session_state
encoded_content = event.content.model_dump(exclude_none=True)
storage_event = StorageEvent( storage_event = StorageEvent(
id=event.id, id=event.id,
invocation_id=event.invocation_id, invocation_id=event.invocation_id,
author=event.author, author=event.author,
branch=event.branch, branch=event.branch,
content=encoded_content,
actions=event.actions, actions=event.actions,
session_id=session.id, session_id=session.id,
app_name=session.app_name, app_name=session.app_name,
user_id=session.user_id, user_id=session.user_id,
timestamp=datetime.fromtimestamp(event.timestamp), timestamp=datetime.fromtimestamp(event.timestamp),
long_running_tool_ids=event.long_running_tool_ids,
grounding_metadata=event.grounding_metadata,
partial=event.partial,
turn_complete=event.turn_complete,
error_code=event.error_code,
error_message=event.error_message,
interrupted=event.interrupted,
) )
if event.content:
encoded_content = event.content.model_dump(exclude_none=True)
# Workaround for multimodal Content throwing JSON not serializable
# error with SQLAlchemy.
for p in encoded_content["parts"]:
if "inline_data" in p:
p["inline_data"]["data"] = (
base64.b64encode(p["inline_data"]["data"]).decode("utf-8"),
)
storage_event.content = encoded_content
sessionFactory.add(storage_event) sessionFactory.add(storage_event)
@@ -489,8 +536,7 @@ class DatabaseSessionService(BaseSessionService):
user_id: str, user_id: str,
session_id: str, session_id: str,
) -> ListEventsResponse: ) -> ListEventsResponse:
pass raise NotImplementedError()
def convert_event(event: StorageEvent) -> Event: def convert_event(event: StorageEvent) -> Event:
"""Converts a storage event to an event.""" """Converts a storage event to an event."""
@@ -505,7 +551,7 @@ def convert_event(event: StorageEvent) -> Event:
) )
def _extract_state_delta(state: dict): def _extract_state_delta(state: dict[str, Any]):
app_state_delta = {} app_state_delta = {}
user_state_delta = {} user_state_delta = {}
session_state_delta = {} session_state_delta = {}
@@ -528,3 +574,10 @@ def _merge_state(app_state, user_state, session_state):
for key in user_state.keys(): for key in user_state.keys():
merged_state[State.USER_PREFIX + key] = user_state[key] merged_state[State.USER_PREFIX + key] = user_state[key]
return merged_state return merged_state
def _decode_content(content: dict[str, Any]) -> dict[str, Any]:
for p in content["parts"]:
if "inline_data" in p:
p["inline_data"]["data"] = base64.b64decode(p["inline_data"]["data"][0])
return content
@@ -196,11 +196,12 @@ class IntegrationClient:
action_details = connections_client.get_action_schema(action) action_details = connections_client.get_action_schema(action)
input_schema = action_details["inputSchema"] input_schema = action_details["inputSchema"]
output_schema = action_details["outputSchema"] output_schema = action_details["outputSchema"]
action_display_name = action_details["displayName"] # Remove spaces from the display name to generate valid spec
action_display_name = action_details["displayName"].replace(" ", "")
operation = "EXECUTE_ACTION" operation = "EXECUTE_ACTION"
if action == "ExecuteCustomQuery": if action == "ExecuteCustomQuery":
connector_spec["components"]["schemas"][ connector_spec["components"]["schemas"][
f"{action}_Request" f"{action_display_name}_Request"
] = connections_client.execute_custom_query_request() ] = connections_client.execute_custom_query_request()
operation = "EXECUTE_QUERY" operation = "EXECUTE_QUERY"
else: else:
@@ -291,7 +291,7 @@ def _parse_schema_from_parameter(
return schema return schema
raise ValueError( raise ValueError(
f'Failed to parse the parameter {param} of function {func_name} for' f'Failed to parse the parameter {param} of function {func_name} for'
' automatic function calling.Automatic function calling works best with' ' automatic function calling. Automatic function calling works best with'
' simpler function signature schema,consider manually parse your' ' simpler function signature schema,consider manually parse your'
f' function declaration for function {func_name}.' f' function declaration for function {func_name}.'
) )
@@ -11,4 +11,77 @@
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and # See the License for the specific language governing permissions and
# limitations under the License. # limitations under the License.
from .google_api_tool_sets import calendar_tool_set __all__ = [
'bigquery_tool_set',
'calendar_tool_set',
'gmail_tool_set',
'youtube_tool_set',
'slides_tool_set',
'sheets_tool_set',
'docs_tool_set',
]
# Nothing is imported here automatically
# Each tool set will only be imported when accessed
_bigquery_tool_set = None
_calendar_tool_set = None
_gmail_tool_set = None
_youtube_tool_set = None
_slides_tool_set = None
_sheets_tool_set = None
_docs_tool_set = None
def __getattr__(name):
global _bigquery_tool_set, _calendar_tool_set, _gmail_tool_set, _youtube_tool_set, _slides_tool_set, _sheets_tool_set, _docs_tool_set
match name:
case 'bigquery_tool_set':
if _bigquery_tool_set is None:
from .google_api_tool_sets import bigquery_tool_set as bigquery
_bigquery_tool_set = bigquery
return _bigquery_tool_set
case 'calendar_tool_set':
if _calendar_tool_set is None:
from .google_api_tool_sets import calendar_tool_set as calendar
_calendar_tool_set = calendar
return _calendar_tool_set
case 'gmail_tool_set':
if _gmail_tool_set is None:
from .google_api_tool_sets import gmail_tool_set as gmail
_gmail_tool_set = gmail
return _gmail_tool_set
case 'youtube_tool_set':
if _youtube_tool_set is None:
from .google_api_tool_sets import youtube_tool_set as youtube
_youtube_tool_set = youtube
return _youtube_tool_set
case 'slides_tool_set':
if _slides_tool_set is None:
from .google_api_tool_sets import slides_tool_set as slides
_slides_tool_set = slides
return _slides_tool_set
case 'sheets_tool_set':
if _sheets_tool_set is None:
from .google_api_tool_sets import sheets_tool_set as sheets
_sheets_tool_set = sheets
return _sheets_tool_set
case 'docs_tool_set':
if _docs_tool_set is None:
from .google_api_tool_sets import docs_tool_set as docs
_docs_tool_set = docs
return _docs_tool_set
@@ -19,37 +19,94 @@ from .google_api_tool_set import GoogleApiToolSet
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
calendar_tool_set = GoogleApiToolSet.load_tool_set( _bigquery_tool_set = None
api_name="calendar", _calendar_tool_set = None
api_version="v3", _gmail_tool_set = None
) _youtube_tool_set = None
_slides_tool_set = None
_sheets_tool_set = None
_docs_tool_set = None
bigquery_tool_set = GoogleApiToolSet.load_tool_set(
api_name="bigquery",
api_version="v2",
)
gmail_tool_set = GoogleApiToolSet.load_tool_set( def __getattr__(name):
api_name="gmail", """This method dynamically loads and returns GoogleApiToolSet instances for
api_version="v1",
)
youtube_tool_set = GoogleApiToolSet.load_tool_set( various Google APIs. It uses a lazy loading approach, initializing each
api_name="youtube", tool set only when it is first requested. This avoids unnecessary loading
api_version="v3", of tool sets that are not used in a given session.
)
slides_tool_set = GoogleApiToolSet.load_tool_set( Args:
api_name="slides", name (str): The name of the tool set to retrieve (e.g.,
api_version="v1", "bigquery_tool_set").
)
sheets_tool_set = GoogleApiToolSet.load_tool_set( Returns:
api_name="sheets", GoogleApiToolSet: The requested tool set instance.
api_version="v4",
)
docs_tool_set = GoogleApiToolSet.load_tool_set( Raises:
api_name="docs", AttributeError: If the requested tool set name is not recognized.
api_version="v1", """
) global _bigquery_tool_set, _calendar_tool_set, _gmail_tool_set, _youtube_tool_set, _slides_tool_set, _sheets_tool_set, _docs_tool_set
match name:
case "bigquery_tool_set":
if _bigquery_tool_set is None:
_bigquery_tool_set = GoogleApiToolSet.load_tool_set(
api_name="bigquery",
api_version="v2",
)
return _bigquery_tool_set
case "calendar_tool_set":
if _calendar_tool_set is None:
_calendar_tool_set = GoogleApiToolSet.load_tool_set(
api_name="calendar",
api_version="v3",
)
return _calendar_tool_set
case "gmail_tool_set":
if _gmail_tool_set is None:
_gmail_tool_set = GoogleApiToolSet.load_tool_set(
api_name="gmail",
api_version="v1",
)
return _gmail_tool_set
case "youtube_tool_set":
if _youtube_tool_set is None:
_youtube_tool_set = GoogleApiToolSet.load_tool_set(
api_name="youtube",
api_version="v3",
)
return _youtube_tool_set
case "slides_tool_set":
if _slides_tool_set is None:
_slides_tool_set = GoogleApiToolSet.load_tool_set(
api_name="slides",
api_version="v1",
)
return _slides_tool_set
case "sheets_tool_set":
if _sheets_tool_set is None:
_sheets_tool_set = GoogleApiToolSet.load_tool_set(
api_name="sheets",
api_version="v4",
)
return _sheets_tool_set
case "docs_tool_set":
if _docs_tool_set is None:
_docs_tool_set = GoogleApiToolSet.load_tool_set(
api_name="docs",
api_version="v1",
)
return _docs_tool_set
@@ -311,7 +311,9 @@ class GoogleApiToOpenApiConverter:
# Determine the actual endpoint path # Determine the actual endpoint path
# Google often has the format something like 'users.messages.list' # Google often has the format something like 'users.messages.list'
rest_path = method_data.get("path", "/") # flatPath is preferred as it provides the actual path, while path
# might contain variables like {+projectId}
rest_path = method_data.get("flatPath", method_data.get("path", "/"))
if not rest_path.startswith("/"): if not rest_path.startswith("/"):
rest_path = "/" + rest_path rest_path = "/" + rest_path
+25 -2
View File
@@ -16,18 +16,26 @@ from __future__ import annotations
from typing import TYPE_CHECKING from typing import TYPE_CHECKING
from google.genai import types
from typing_extensions import override from typing_extensions import override
from .function_tool import FunctionTool from .function_tool import FunctionTool
from .tool_context import ToolContext from .tool_context import ToolContext
if TYPE_CHECKING: if TYPE_CHECKING:
from ..models import LlmRequest
from ..memory.base_memory_service import MemoryResult from ..memory.base_memory_service import MemoryResult
from ..models import LlmRequest
def load_memory(query: str, tool_context: ToolContext) -> 'list[MemoryResult]': def load_memory(query: str, tool_context: ToolContext) -> 'list[MemoryResult]':
"""Loads the memory for the current user.""" """Loads the memory for the current user.
Args:
query: The query to load the memory for.
Returns:
A list of memory results.
"""
response = tool_context.search_memory(query) response = tool_context.search_memory(query)
return response.memories return response.memories
@@ -38,6 +46,21 @@ class LoadMemoryTool(FunctionTool):
def __init__(self): def __init__(self):
super().__init__(load_memory) super().__init__(load_memory)
@override
def _get_declaration(self) -> types.FunctionDeclaration | None:
return types.FunctionDeclaration(
name=self.name,
description=self.description,
parameters=types.Schema(
type=types.Type.OBJECT,
properties={
'query': types.Schema(
type=types.Type.STRING,
)
},
),
)
@override @override
async def process_llm_request( async def process_llm_request(
self, self,

Some files were not shown because too many files have changed in this diff Show More