Skip to content

Commit ac968b5

Browse files
committed
feat: integrate API key bridging and enhance checkpointer functionality
- Introduced `bridge_provider_api_keys` function to map MISAKA_* API keys into standard environment variables for LangChain, ensuring existing keys are not overwritten. - Enhanced checkpointer setup and teardown processes, improving error handling and logging during initialization and closure. - Updated main application lifecycle to include API key bridging during startup. - Added tests for the new API key bridging functionality to ensure correct behavior and environment variable management.
1 parent b115d1c commit ac968b5

24 files changed

Lines changed: 1517 additions & 135 deletions

agent/app/config.py

Lines changed: 22 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -2,12 +2,16 @@
22

33
from __future__ import annotations
44

5+
import logging
6+
import os
57
from functools import lru_cache
68
from pathlib import Path
79

810
from pydantic import Field
911
from pydantic_settings import BaseSettings
1012

13+
logger = logging.getLogger(__name__)
14+
1115

1216
def _default_interrupt_config() -> dict[str, bool]:
1317
return {
@@ -57,4 +61,22 @@ def get_settings() -> Settings:
5761
return Settings()
5862

5963

64+
def bridge_provider_api_keys(settings_obj: Settings | None = None) -> None:
65+
"""Bridge MISAKA_* keys into standard provider env vars used by LangChain.
66+
67+
Does not overwrite keys that are already present in the process environment.
68+
Never logs key values.
69+
"""
70+
cfg = settings_obj or get_settings()
71+
bridged: list[str] = []
72+
if cfg.anthropic_api_key and not os.environ.get("ANTHROPIC_API_KEY"):
73+
os.environ["ANTHROPIC_API_KEY"] = cfg.anthropic_api_key
74+
bridged.append("ANTHROPIC_API_KEY")
75+
if cfg.openai_api_key and not os.environ.get("OPENAI_API_KEY"):
76+
os.environ["OPENAI_API_KEY"] = cfg.openai_api_key
77+
bridged.append("OPENAI_API_KEY")
78+
if bridged:
79+
logger.info("Bridged provider API keys into process env: %s", ", ".join(bridged))
80+
81+
6082
settings = Settings()

agent/app/dependencies.py

Lines changed: 34 additions & 33 deletions
Original file line numberDiff line numberDiff line change
@@ -11,6 +11,7 @@
1111
logger = logging.getLogger(__name__)
1212

1313
_checkpointer: Any | None = None
14+
_checkpointer_cm: Any | None = None
1415
_store: Any | None = None
1516

1617

@@ -32,10 +33,15 @@ def get_store():
3233

3334

3435
def get_checkpointer():
35-
"""Return the AsyncSqliteSaver singleton, creating it lazily."""
36-
global _checkpointer
36+
"""Return the process-wide AsyncSqliteSaver, or None if unavailable/uninitialized."""
37+
return _checkpointer
38+
39+
40+
async def setup_checkpointer() -> None:
41+
"""Enter AsyncSqliteSaver context, create schema, and enable WAL via setup()."""
42+
global _checkpointer, _checkpointer_cm
3743
if _checkpointer is not None:
38-
return _checkpointer
44+
return
3945

4046
settings = get_settings()
4147
settings.data_dir.mkdir(parents=True, exist_ok=True)
@@ -44,40 +50,35 @@ def get_checkpointer():
4450
from langgraph.checkpoint.sqlite.aio import AsyncSqliteSaver
4551
except ImportError:
4652
logger.warning("AsyncSqliteSaver unavailable; continuing without checkpointer")
47-
return None
53+
return
4854

4955
conn_string = str(settings.checkpointer_db_path)
50-
_checkpointer = AsyncSqliteSaver.from_conn_string(conn_string)
51-
return _checkpointer
52-
53-
54-
async def setup_checkpointer() -> None:
55-
"""Ensure checkpointer schema exists."""
56-
checkpointer = get_checkpointer()
57-
if checkpointer is None:
58-
return
59-
setup = getattr(checkpointer, "setup", None)
60-
if setup is None:
56+
cm = AsyncSqliteSaver.from_conn_string(conn_string)
57+
try:
58+
saver = await cm.__aenter__()
59+
await saver.setup()
60+
except Exception:
61+
logger.exception("Failed to initialize checkpointer")
62+
try:
63+
await cm.__aexit__(None, None, None)
64+
except Exception:
65+
logger.exception("Failed to clean up checkpointer after setup error")
6166
return
62-
result = setup()
63-
if hasattr(result, "__await__"):
64-
await result
6567

68+
_checkpointer_cm = cm
69+
_checkpointer = saver
70+
logger.info("Checkpointer ready at %s (WAL enabled via setup)", conn_string)
6671

67-
async def close_checkpointer() -> None:
68-
"""Close checkpointer resources if present."""
69-
global _checkpointer
70-
checkpointer = _checkpointer
71-
if checkpointer is None:
72-
return
73-
74-
for method_name in ("aclose", "close"):
75-
method = getattr(checkpointer, method_name, None)
76-
if method is None:
77-
continue
78-
result = method()
79-
if hasattr(result, "__await__"):
80-
await result
81-
break
8272

73+
async def close_checkpointer() -> None:
74+
"""Exit the AsyncSqliteSaver context manager and clear the singleton."""
75+
global _checkpointer, _checkpointer_cm
76+
cm = _checkpointer_cm
8377
_checkpointer = None
78+
_checkpointer_cm = None
79+
if cm is None:
80+
return
81+
try:
82+
await cm.__aexit__(None, None, None)
83+
except Exception:
84+
logger.exception("Failed to close checkpointer")

agent/app/main.py

Lines changed: 4 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -5,17 +5,19 @@
55

66
from fastapi import FastAPI
77

8-
from app.config import settings
8+
from app.config import bridge_provider_api_keys, settings
99
from app.dependencies import close_checkpointer, setup_checkpointer
1010
from app.routers.agent import router as agent_router
1111
from app.routers.health import router as health_router
1212
from app.routers.info import router as info_router
13+
from app.routers.memory import router as memory_router
1314

1415

1516
@asynccontextmanager
1617
async def lifespan(application: FastAPI):
1718
"""Manage startup and shutdown lifecycle."""
1819
application.state.startup_time = time.time()
20+
bridge_provider_api_keys()
1921
await setup_checkpointer()
2022
try:
2123
yield
@@ -34,3 +36,4 @@ async def lifespan(application: FastAPI):
3436
app.include_router(health_router)
3537
app.include_router(info_router)
3638
app.include_router(agent_router)
39+
app.include_router(memory_router)

agent/app/routers/memory.py

Lines changed: 149 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,149 @@
1+
"""PowerMem REST endpoints for search / list / delete."""
2+
3+
from __future__ import annotations
4+
5+
import logging
6+
from typing import Any
7+
8+
from fastapi import APIRouter, HTTPException, Query
9+
from pydantic import BaseModel, Field
10+
11+
from app.memory import get_memory_engine
12+
13+
router = APIRouter(prefix="/memory", tags=["memory"])
14+
logger = logging.getLogger(__name__)
15+
16+
17+
class MemoryItem(BaseModel):
18+
id: str | None = None
19+
content: str = ""
20+
score: float | None = None
21+
created_at: str | None = None
22+
metadata: dict[str, Any] = Field(default_factory=dict)
23+
24+
25+
class MemoryListResponse(BaseModel):
26+
items: list[MemoryItem]
27+
offset: int = 0
28+
limit: int = 20
29+
30+
31+
def _require_engine() -> Any:
32+
engine = get_memory_engine()
33+
if engine is None:
34+
raise HTTPException(status_code=503, detail="Memory engine unavailable")
35+
return engine
36+
37+
38+
def _normalize_item(item: Any) -> MemoryItem:
39+
if isinstance(item, dict):
40+
return MemoryItem(
41+
id=_as_optional_str(item.get("id") or item.get("memory_id")),
42+
content=str(item.get("content") or item.get("memory") or ""),
43+
score=_as_optional_float(item.get("score")),
44+
created_at=_as_optional_str(item.get("created_at")),
45+
metadata=item.get("metadata") if isinstance(item.get("metadata"), dict) else {},
46+
)
47+
48+
return MemoryItem(
49+
id=_as_optional_str(getattr(item, "id", None) or getattr(item, "memory_id", None)),
50+
content=str(
51+
getattr(item, "content", None) or getattr(item, "memory", None) or ""
52+
),
53+
score=_as_optional_float(getattr(item, "score", None)),
54+
created_at=_as_optional_str(getattr(item, "created_at", None)),
55+
metadata=getattr(item, "metadata", None)
56+
if isinstance(getattr(item, "metadata", None), dict)
57+
else {},
58+
)
59+
60+
61+
def _as_optional_str(value: Any) -> str | None:
62+
if value is None:
63+
return None
64+
return str(value)
65+
66+
67+
def _as_optional_float(value: Any) -> float | None:
68+
if value is None:
69+
return None
70+
try:
71+
return float(value)
72+
except (TypeError, ValueError):
73+
return None
74+
75+
76+
def _call_search(engine: Any, query: str, limit: int) -> list[Any]:
77+
if hasattr(engine, "search"):
78+
result = engine.search(query, limit=limit)
79+
elif hasattr(engine, "query"):
80+
result = engine.query(query, limit=limit)
81+
else:
82+
raise HTTPException(status_code=503, detail="Memory engine has no search API")
83+
return list(result or [])
84+
85+
86+
def _call_list(engine: Any, offset: int, limit: int) -> list[Any]:
87+
if hasattr(engine, "list"):
88+
result = engine.list(offset=offset, limit=limit)
89+
elif hasattr(engine, "get_all"):
90+
result = engine.get_all()
91+
result = list(result or [])[offset : offset + limit]
92+
else:
93+
raise HTTPException(status_code=503, detail="Memory engine has no list API")
94+
return list(result or [])
95+
96+
97+
def _call_delete(engine: Any, memory_id: str) -> None:
98+
if hasattr(engine, "delete"):
99+
engine.delete(memory_id)
100+
return
101+
if hasattr(engine, "remove"):
102+
engine.remove(memory_id)
103+
return
104+
raise HTTPException(status_code=503, detail="Memory engine has no delete API")
105+
106+
107+
@router.get("/search", response_model=MemoryListResponse)
108+
async def memory_search(
109+
query: str = Query(..., min_length=1),
110+
limit: int = Query(5, ge=1, le=100),
111+
) -> MemoryListResponse:
112+
engine = _require_engine()
113+
try:
114+
items = [_normalize_item(item) for item in _call_search(engine, query, limit)]
115+
except HTTPException:
116+
raise
117+
except Exception as exc:
118+
logger.exception("memory search failed")
119+
raise HTTPException(status_code=500, detail=str(exc)) from exc
120+
return MemoryListResponse(items=items, offset=0, limit=limit)
121+
122+
123+
@router.get("/list", response_model=MemoryListResponse)
124+
async def memory_list(
125+
offset: int = Query(0, ge=0),
126+
limit: int = Query(20, ge=1, le=100),
127+
) -> MemoryListResponse:
128+
engine = _require_engine()
129+
try:
130+
items = [_normalize_item(item) for item in _call_list(engine, offset, limit)]
131+
except HTTPException:
132+
raise
133+
except Exception as exc:
134+
logger.exception("memory list failed")
135+
raise HTTPException(status_code=500, detail=str(exc)) from exc
136+
return MemoryListResponse(items=items, offset=offset, limit=limit)
137+
138+
139+
@router.delete("/{memory_id}")
140+
async def memory_delete(memory_id: str) -> dict[str, str]:
141+
engine = _require_engine()
142+
try:
143+
_call_delete(engine, memory_id)
144+
except HTTPException:
145+
raise
146+
except Exception as exc:
147+
logger.exception("memory delete failed")
148+
raise HTTPException(status_code=500, detail=str(exc)) from exc
149+
return {"status": "deleted", "id": memory_id}

agent/tests/test_config.py

Lines changed: 19 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -59,3 +59,22 @@ def test_get_settings_returns_singleton():
5959
second = get_settings()
6060
assert first is second
6161
get_settings.cache_clear()
62+
63+
64+
def test_bridge_provider_api_keys(monkeypatch):
65+
from app.config import Settings, bridge_provider_api_keys
66+
67+
monkeypatch.delenv("ANTHROPIC_API_KEY", raising=False)
68+
monkeypatch.delenv("OPENAI_API_KEY", raising=False)
69+
cfg = Settings(
70+
anthropic_api_key="sk-ant-test",
71+
openai_api_key="sk-openai-test",
72+
)
73+
bridge_provider_api_keys(cfg)
74+
assert __import__("os").environ["ANTHROPIC_API_KEY"] == "sk-ant-test"
75+
assert __import__("os").environ["OPENAI_API_KEY"] == "sk-openai-test"
76+
77+
# Does not overwrite existing values.
78+
monkeypatch.setenv("ANTHROPIC_API_KEY", "keep-me")
79+
bridge_provider_api_keys(cfg)
80+
assert __import__("os").environ["ANTHROPIC_API_KEY"] == "keep-me"

0 commit comments

Comments
 (0)