diff --git a/server/routers/entities.py b/server/routers/entities.py index 5a5a0a42b..d7d2875c6 100644 --- a/server/routers/entities.py +++ b/server/routers/entities.py @@ -5,7 +5,7 @@ from typing import Any, Literal, Optional from fastapi import APIRouter, Depends from pydantic import BaseModel -from auth import verify_auth +from auth import require_admin from errors import upstream_error from schemas import MessageResponse from server_state import get_memory_instance @@ -42,7 +42,7 @@ def _parse_timestamp(value: Any) -> Optional[datetime]: @router.get("", response_model=list[Entity]) -def list_entities(_auth=Depends(verify_auth)): +def list_entities(_admin=Depends(require_admin)): buckets: dict[tuple[EntityType, str], dict[str, Any]] = defaultdict( lambda: {"total_memories": 0, "created_at": None, "updated_at": None} ) @@ -69,7 +69,7 @@ def list_entities(_auth=Depends(verify_auth)): @router.delete("/{entity_type}/{entity_id}", response_model=MessageResponse) -def delete_entity(entity_type: EntityType, entity_id: str, _auth=Depends(verify_auth)): +def delete_entity(entity_type: EntityType, entity_id: str, _admin=Depends(require_admin)): try: get_memory_instance().delete_all(**{TYPE_TO_FIELD[entity_type]: entity_id}) except Exception: diff --git a/server/tests/test_admin_gating.py b/server/tests/test_admin_gating.py index 7c08d4f34..107aac1b2 100644 --- a/server/tests/test_admin_gating.py +++ b/server/tests/test_admin_gating.py @@ -101,3 +101,59 @@ def test_post_reset_member_forbidden(client, auth_member_header): def test_post_reset_no_auth_unauthorized(client): response = client.post("/reset") assert response.status_code == 401 + + +# --- GET /entities --- + + +def test_get_entities_admin_jwt(client, auth_admin_header): + response = client.get("/entities", headers=auth_admin_header) + assert response.status_code == 200 + + +def test_get_entities_admin_api_key(client, admin_api_key_env): + response = client.get("/entities", headers=admin_api_key_env) + assert response.status_code == 200 + + +def test_get_entities_auth_disabled(client, auth_disabled_env): + response = client.get("/entities") + assert response.status_code == 200 + + +def test_get_entities_member_forbidden(client, auth_member_header): + response = client.get("/entities", headers=auth_member_header) + assert response.status_code == 403 + + +def test_get_entities_no_auth_unauthorized(client): + response = client.get("/entities") + assert response.status_code == 401 + + +# --- DELETE /entities/{type}/{id} --- + + +def test_delete_entity_admin_jwt(client, auth_admin_header): + response = client.delete("/entities/user/alice", headers=auth_admin_header) + assert response.status_code == 200 + + +def test_delete_entity_admin_api_key(client, admin_api_key_env): + response = client.delete("/entities/user/alice", headers=admin_api_key_env) + assert response.status_code == 200 + + +def test_delete_entity_auth_disabled(client, auth_disabled_env): + response = client.delete("/entities/user/alice") + assert response.status_code == 200 + + +def test_delete_entity_member_forbidden(client, auth_member_header): + response = client.delete("/entities/user/alice", headers=auth_member_header) + assert response.status_code == 403 + + +def test_delete_entity_no_auth_unauthorized(client): + response = client.delete("/entities/user/alice") + assert response.status_code == 401