diff --git a/pyproject.toml b/pyproject.toml index 3b4ef0be..a646747c 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -51,6 +51,7 @@ dependencies = [ "typing-extensions>=4.5, <5", "tzlocal>=5.3", # Time zone utilities "uvicorn>=0.34.0", # ASGI server for FastAPI + "watchdog>=6.0.0", # For file change detection and hot reload "websockets>=15.0.1", # For BaseLlmFlow # go/keep-sorted end ] diff --git a/src/google/adk/cli/cli_tools_click.py b/src/google/adk/cli/cli_tools_click.py index d9886dc8..95e61778 100644 --- a/src/google/adk/cli/cli_tools_click.py +++ b/src/google/adk/cli/cli_tools_click.py @@ -583,6 +583,13 @@ def fast_api_common_options(): default=False, help="Optional. Whether to enable A2A endpoint.", ) + @click.option( + "--reload_agents", + is_flag=True, + default=False, + show_default=True, + help="Optional. Whether to enable live reload for agents changes.", + ) @functools.wraps(func) def wrapper(*args, **kwargs): return func(*args, **kwargs) @@ -625,6 +632,7 @@ def cli_web( session_db_url: Optional[str] = None, # Deprecated artifact_storage_uri: Optional[str] = None, # Deprecated a2a: bool = False, + reload_agents: bool = False, ): """Starts a FastAPI server with Web UI for agents. @@ -674,6 +682,7 @@ def cli_web( a2a=a2a, host=host, port=port, + reload_agents=reload_agents, ) config = uvicorn.Config( app, @@ -721,6 +730,7 @@ def cli_api_server( session_db_url: Optional[str] = None, # Deprecated artifact_storage_uri: Optional[str] = None, # Deprecated a2a: bool = False, + reload_agents: bool = False, ): """Starts a FastAPI server for agents. @@ -748,6 +758,7 @@ def cli_api_server( a2a=a2a, host=host, port=port, + reload_agents=reload_agents, ), host=host, port=port, @@ -861,6 +872,7 @@ def cli_deploy_cloud_run( session_db_url: Optional[str] = None, # Deprecated artifact_storage_uri: Optional[str] = None, # Deprecated a2a: bool = False, + reload_agents: bool = False, ): """Deploys an agent to Cloud Run. diff --git a/src/google/adk/cli/fast_api.py b/src/google/adk/cli/fast_api.py index 13605147..83cd8999 100644 --- a/src/google/adk/cli/fast_api.py +++ b/src/google/adk/cli/fast_api.py @@ -49,6 +49,8 @@ from pydantic import Field from pydantic import ValidationError from starlette.types import Lifespan from typing_extensions import override +from watchdog.events import FileSystemEventHandler +from watchdog.observers import Observer from ..agents import RunConfig from ..agents.live_request_queue import LiveRequest @@ -87,6 +89,21 @@ from .utils.agent_loader import AgentLoader logger = logging.getLogger("google_adk." + __name__) _EVAL_SET_FILE_EXTENSION = ".evalset.json" +_app_name = "" +_runners_to_clean = set() + + +class AgentChangeEventHandler(FileSystemEventHandler): + + def __init__(self, agent_loader: AgentLoader): + self.agent_loader = agent_loader + + def on_modified(self, event): + if not (event.src_path.endswith(".py") or event.src_path.endswith(".yaml")): + return + logger.info("Change detected in agents directory: %s", event.src_path) + self.agent_loader.remove_agent_from_cache(_app_name) + _runners_to_clean.add(_app_name) class ApiServerSpanExporter(export.SpanExporter): @@ -205,6 +222,7 @@ def get_fast_api_app( host: str = "127.0.0.1", port: int = 8000, trace_to_cloud: bool = False, + reload_agents: bool = False, lifespan: Optional[Lifespan[FastAPI]] = None, ) -> FastAPI: # InMemory tracing dict. @@ -235,7 +253,6 @@ def get_fast_api_app( @asynccontextmanager async def internal_lifespan(app: FastAPI): - try: if lifespan: async with lifespan(app) as lifespan_context: @@ -243,6 +260,9 @@ def get_fast_api_app( else: yield finally: + if reload_agents: + observer.stop() + observer.join() # Create tasks for all runner closures to run concurrently await cleanup.close_runners(list(runner_dict.values())) @@ -336,6 +356,13 @@ def get_fast_api_app( # initialize Agent Loader agent_loader = AgentLoader(agents_dir) + # Set up a file system watcher to detect changes in the agents directory. + observer = Observer() + if reload_agents: + event_handler = AgentChangeEventHandler(agent_loader) + observer.schedule(event_handler, agents_dir, recursive=True) + observer.start() + @app.get("/list-apps") def list_apps() -> list[str]: base_path = Path.cwd() / agents_dir @@ -390,6 +417,9 @@ def get_fast_api_app( ) if not session: raise HTTPException(status_code=404, detail="Session not found") + + global _app_name + _app_name = app_name return session @app.get( @@ -947,6 +977,11 @@ def get_fast_api_app( async def _get_runner_async(app_name: str) -> Runner: """Returns the runner for the given app.""" + if app_name in _runners_to_clean: + _runners_to_clean.remove(app_name) + runner = runner_dict.pop(app_name, None) + await cleanup.close_runners(list([runner])) + envs.load_dotenv_for_agent(os.path.basename(app_name), agents_dir) if app_name in runner_dict: return runner_dict[app_name] diff --git a/src/google/adk/cli/utils/agent_loader.py b/src/google/adk/cli/utils/agent_loader.py index cd61dfbf..ca81bd23 100644 --- a/src/google/adk/cli/utils/agent_loader.py +++ b/src/google/adk/cli/utils/agent_loader.py @@ -164,3 +164,15 @@ class AgentLoader: agent = self._perform_load(agent_name) self._agent_cache[agent_name] = agent return agent + + def remove_agent_from_cache(self, agent_name: str): + # Clear module cache for the agent and its submodules + keys_to_delete = [ + module_name + for module_name in sys.modules + if module_name == agent_name or module_name.startswith(f"{agent_name}.") + ] + for key in keys_to_delete: + logger.debug("Deleting module %s", key) + del sys.modules[key] + self._agent_cache.pop(agent_name, None)