diff --git a/src/hai_agents/client.py b/src/hai_agents/client.py index d53c090..c00d0a9 100644 --- a/src/hai_agents/client.py +++ b/src/hai_agents/client.py @@ -59,13 +59,18 @@ def _stored_api_key() -> typing.Optional[str]: return None +def default_api_key() -> typing.Optional[str]: + """`HAI_API_KEY`, else the key stored by `hai login`.""" + return os.getenv(API_KEY_VAR) or _stored_api_key() + + def _default_api_key(init: typing.Callable[_P, None]) -> typing.Callable[_P, None]: """Resolve `api_key` as: argument, then `HAI_API_KEY`, then the key stored by `hai login`.""" @functools.wraps(init) def wrapper(*args: _P.args, **kwargs: _P.kwargs) -> None: if kwargs.get("api_key") is None: - api_key = os.getenv(API_KEY_VAR) or _stored_api_key() + api_key = default_api_key() if api_key is None: raise ApiError(body=f"No API key found. Pass api_key, set {API_KEY_VAR}, or run `hai login`.") kwargs["api_key"] = api_key diff --git a/src/hai_agents_local/runtime/acquire.py b/src/hai_agents_local/runtime/acquire.py index 5c7508b..3c4a0a8 100644 --- a/src/hai_agents_local/runtime/acquire.py +++ b/src/hai_agents_local/runtime/acquire.py @@ -6,6 +6,8 @@ import threading import typing +from hai_agents.client import API_KEY_VAR, default_api_key + from .inference import Inference from .runtime import LocalRuntime @@ -61,6 +63,9 @@ def _launch_options( options = dict(local_options or {}) options["required_recipe"] = SHARED_RECIPE options["spawn_env"] = {RECIPE_ENV: SHARED_RECIPE, **options.get("spawn_env", {})} + api_key = default_api_key() + if api_key is not None: + options["spawn_env"].setdefault(API_KEY_VAR, api_key) if inference is not None: options["spawn_env"] = inference.runtime_env(options["spawn_env"]) options["inherit_env"] = False diff --git a/tests/test_runtime_placement.py b/tests/test_runtime_placement.py index 5b8294a..5344f5e 100644 --- a/tests/test_runtime_placement.py +++ b/tests/test_runtime_placement.py @@ -161,6 +161,26 @@ def test_self_hosted_inference_does_not_receive_hosted_key(monkeypatch): assert "HAI_AGENT_RUNTIME_BASE_URL" not in Inference.cloud().runtime_env() +@pytest.mark.parametrize( + ("inference", "forwarded"), + [(None, "hk-stored"), (Inference.self_hosted("http://127.0.0.1:8000/v1", model="my-model"), None)], +) +def test_spawned_runtime_infers_with_the_key_hai_login_stored(tmp_path, monkeypatch, inference, forwarded): + monkeypatch.delenv("HAI_API_KEY", raising=False) + monkeypatch.setenv("XDG_CONFIG_HOME", str(tmp_path)) + (tmp_path / "hai").mkdir() + (tmp_path / "hai" / ".env").write_text("HAI_API_KEY=hk-stored\n") + launches = [] + + def start(**options): + launches.append(options) + return FakeRuntime() + + monkeypatch.setattr(LocalRuntime, "ensure_started", start) + Client.local(inference=inference, auto_bridges=False) + assert launches[0]["spawn_env"].get("HAI_API_KEY") == forwarded + + @pytest.mark.parametrize("proof_token", ["local-token", "squatter-token", None]) def test_only_the_runtime_holding_the_token_ever_receives_it(tmp_path, runtime_server, proof_token): write_owner_only(token_file_path(runtime_server.port, cache_dir=tmp_path), "local-token")