diff --git a/server/main.py b/server/main.py index 8adc3c74e..89683b3dd 100644 --- a/server/main.py +++ b/server/main.py @@ -13,7 +13,7 @@ from slowapi import _rate_limit_exceeded_handler from slowapi.errors import RateLimitExceeded from sqlalchemy import func, select -from auth import ADMIN_API_KEY, AUTH_DISABLED, JWT_SECRET, require_admin, verify_auth +from auth import ADMIN_API_KEY, AUTH_DISABLED, JWT_SECRET, _ensure_admin, require_admin, verify_auth from errors import ( UpstreamError, install_request_id_logging, @@ -386,19 +386,27 @@ def _list_all_memories(limit: int = ALL_MEMORIES_LIMIT) -> Dict[str, Any]: @app.get("/memories", summary="Get memories") def get_all_memories( + request: Request, user_id: Optional[str] = None, run_id: Optional[str] = None, agent_id: Optional[str] = None, - _auth=Depends(verify_auth), + user: User | None = Depends(verify_auth), ): - """Retrieve stored memories. Lists all memories when no identifier is provided.""" + """Retrieve stored memories. Lists all memories when no identifier is provided. + + Note: the unfiltered listing branch is admin-only — it would otherwise leak + every payload in the vector store across tenants once multi-user lands. + Filtered queries remain available to any authenticated caller.""" try: if not any([user_id, run_id, agent_id]): + _ensure_admin(request, user) return _list_all_memories() filters = { k: v for k, v in {"user_id": user_id, "run_id": run_id, "agent_id": agent_id}.items() if v is not None } return get_memory_instance().get_all(filters=filters) + except HTTPException: + raise except Exception: raise upstream_error() diff --git a/server/tests/test_admin_gating.py b/server/tests/test_admin_gating.py index 0d30a68b2..e236b851c 100644 --- a/server/tests/test_admin_gating.py +++ b/server/tests/test_admin_gating.py @@ -185,3 +185,38 @@ def test_get_requests_member_forbidden(client, auth_member_header): def test_get_requests_no_auth_unauthorized(client): response = client.get("/requests") assert response.status_code == 401 + + +# --- GET /memories (no filter branch only) --- + + +def test_get_memories_no_filter_admin_jwt(client, auth_admin_header): + response = client.get("/memories", headers=auth_admin_header) + assert response.status_code == 200 + + +def test_get_memories_no_filter_admin_api_key(client, admin_api_key_env): + response = client.get("/memories", headers=admin_api_key_env) + assert response.status_code == 200 + + +def test_get_memories_no_filter_auth_disabled(client, auth_disabled_env): + response = client.get("/memories") + assert response.status_code == 200 + + +def test_get_memories_no_filter_member_forbidden(client, auth_member_header): + """The info-disclosure branch (_list_all_memories) is admin-only.""" + response = client.get("/memories", headers=auth_member_header) + assert response.status_code == 403 + + +def test_get_memories_filtered_member_succeeds(client, auth_member_header): + """Filtered queries are unchanged — member can still call them.""" + response = client.get("/memories?user_id=alice", headers=auth_member_header) + assert response.status_code == 200 + + +def test_get_memories_no_auth_unauthorized(client): + response = client.get("/memories") + assert response.status_code == 401