fix: address PR review feedback (round 5)

- Validate OCS envelope in trash_collective, delete_collective, trash_page
- Guard _unwrap_ocs against non-OCS responses with informative OCSError
- Remove _get_ocs_headers() indirection, use class constants directly
- Split headers: _OCS_HEADERS (GET) vs _OCS_HEADERS_JSON (with body)
- Fix docstring claiming emoji param is required when it is optional
- Rename misleading test, add test for non-OCS envelope handling

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
Chris Coutinho
2026-03-26 13:49:28 +01:00
co-authored by Claude Opus 4.6
parent 95edd9ba8e
commit aa46c6147b
3 changed files with 52 additions and 34 deletions
+31 -26
View File
@@ -26,17 +26,19 @@ class CollectivesClient(BaseNextcloudClient):
_OCS_HEADERS: dict[str, str] = { _OCS_HEADERS: dict[str, str] = {
"OCS-APIRequest": "true", "OCS-APIRequest": "true",
"Content-Type": "application/json",
"Accept": "application/json", "Accept": "application/json",
} }
def _get_ocs_headers(self) -> dict[str, str]: _OCS_HEADERS_JSON: dict[str, str] = {
"""Get standard headers required for OCS API calls.""" **_OCS_HEADERS,
return self._OCS_HEADERS "Content-Type": "application/json",
}
def _unwrap_ocs(self, response_json: dict[str, Any]) -> Any: def _unwrap_ocs(self, response_json: dict[str, Any]) -> Any:
"""Unwrap OCS envelope, validating the status before returning data.""" """Unwrap OCS envelope, validating the status before returning data."""
ocs = response_json["ocs"] ocs = response_json.get("ocs")
if ocs is None:
raise OCSError(500, "Response is not an OCS envelope")
meta = ocs.get("meta", {}) meta = ocs.get("meta", {})
status_code = meta.get("statuscode", 200) status_code = meta.get("statuscode", 200)
if status_code >= 400: if status_code >= 400:
@@ -49,7 +51,7 @@ class CollectivesClient(BaseNextcloudClient):
async def get_collectives(self) -> list[dict[str, Any]]: async def get_collectives(self) -> list[dict[str, Any]]:
"""List all collectives the user has access to.""" """List all collectives the user has access to."""
response = await self._make_request( response = await self._make_request(
"GET", f"{API_BASE}/collectives", headers=self._get_ocs_headers() "GET", f"{API_BASE}/collectives", headers=self._OCS_HEADERS
) )
data = self._unwrap_ocs(response.json()) data = self._unwrap_ocs(response.json())
return data["collectives"] return data["collectives"]
@@ -65,7 +67,7 @@ class CollectivesClient(BaseNextcloudClient):
"POST", "POST",
f"{API_BASE}/collectives", f"{API_BASE}/collectives",
json=json_data, json=json_data,
headers=self._get_ocs_headers(), headers=self._OCS_HEADERS_JSON,
) )
data = self._unwrap_ocs(response.json()) data = self._unwrap_ocs(response.json())
return data["collective"] return data["collective"]
@@ -87,18 +89,19 @@ class CollectivesClient(BaseNextcloudClient):
"PUT", "PUT",
f"{API_BASE}/collectives/{collective_id}", f"{API_BASE}/collectives/{collective_id}",
json=json_data, json=json_data,
headers=self._get_ocs_headers(), headers=self._OCS_HEADERS_JSON,
) )
data = self._unwrap_ocs(response.json()) data = self._unwrap_ocs(response.json())
return data["collective"] return data["collective"]
async def trash_collective(self, collective_id: int) -> None: async def trash_collective(self, collective_id: int) -> None:
"""Move a collective to trash (soft delete).""" """Move a collective to trash (soft delete)."""
await self._make_request( response = await self._make_request(
"DELETE", "DELETE",
f"{API_BASE}/collectives/{collective_id}", f"{API_BASE}/collectives/{collective_id}",
headers=self._get_ocs_headers(), headers=self._OCS_HEADERS,
) )
self._unwrap_ocs(response.json())
async def delete_collective(self, collective_id: int) -> None: async def delete_collective(self, collective_id: int) -> None:
"""Permanently delete a collective (must be trashed first). """Permanently delete a collective (must be trashed first).
@@ -106,11 +109,12 @@ class CollectivesClient(BaseNextcloudClient):
This is irreversible. The collective must be in the trash before This is irreversible. The collective must be in the trash before
calling this method. calling this method.
""" """
await self._make_request( response = await self._make_request(
"DELETE", "DELETE",
f"{API_BASE}/collectives/trash/{collective_id}", f"{API_BASE}/collectives/trash/{collective_id}",
headers=self._get_ocs_headers(), headers=self._OCS_HEADERS,
) )
self._unwrap_ocs(response.json())
# Pages # Pages
@@ -119,7 +123,7 @@ class CollectivesClient(BaseNextcloudClient):
response = await self._make_request( response = await self._make_request(
"GET", "GET",
f"{API_BASE}/collectives/{collective_id}/pages", f"{API_BASE}/collectives/{collective_id}/pages",
headers=self._get_ocs_headers(), headers=self._OCS_HEADERS,
) )
data = self._unwrap_ocs(response.json()) data = self._unwrap_ocs(response.json())
return data["pages"] return data["pages"]
@@ -129,7 +133,7 @@ class CollectivesClient(BaseNextcloudClient):
response = await self._make_request( response = await self._make_request(
"GET", "GET",
f"{API_BASE}/collectives/{collective_id}/pages/{page_id}", f"{API_BASE}/collectives/{collective_id}/pages/{page_id}",
headers=self._get_ocs_headers(), headers=self._OCS_HEADERS,
) )
data = self._unwrap_ocs(response.json()) data = self._unwrap_ocs(response.json())
return data["page"] return data["page"]
@@ -143,7 +147,7 @@ class CollectivesClient(BaseNextcloudClient):
"POST", "POST",
f"{API_BASE}/collectives/{collective_id}/pages/{parent_id}", f"{API_BASE}/collectives/{collective_id}/pages/{parent_id}",
json=json_data, json=json_data,
headers=self._get_ocs_headers(), headers=self._OCS_HEADERS_JSON,
) )
data = self._unwrap_ocs(response.json()) data = self._unwrap_ocs(response.json())
return data["page"] return data["page"]
@@ -167,18 +171,19 @@ class CollectivesClient(BaseNextcloudClient):
"PUT", "PUT",
f"{API_BASE}/collectives/{collective_id}/pages/{page_id}", f"{API_BASE}/collectives/{collective_id}/pages/{page_id}",
json=json_data, json=json_data,
headers=self._get_ocs_headers(), headers=self._OCS_HEADERS_JSON,
) )
data = self._unwrap_ocs(response.json()) data = self._unwrap_ocs(response.json())
return data["page"] return data["page"]
async def trash_page(self, collective_id: int, page_id: int) -> None: async def trash_page(self, collective_id: int, page_id: int) -> None:
"""Move a page to trash (soft delete).""" """Move a page to trash (soft delete)."""
await self._make_request( response = await self._make_request(
"DELETE", "DELETE",
f"{API_BASE}/collectives/{collective_id}/pages/{page_id}", f"{API_BASE}/collectives/{collective_id}/pages/{page_id}",
headers=self._get_ocs_headers(), headers=self._OCS_HEADERS,
) )
self._unwrap_ocs(response.json())
async def set_page_emoji( async def set_page_emoji(
self, collective_id: int, page_id: int, emoji: str | None self, collective_id: int, page_id: int, emoji: str | None
@@ -189,7 +194,7 @@ class CollectivesClient(BaseNextcloudClient):
"PUT", "PUT",
f"{API_BASE}/collectives/{collective_id}/pages/{page_id}/emoji", f"{API_BASE}/collectives/{collective_id}/pages/{page_id}/emoji",
json=json_data, json=json_data,
headers=self._get_ocs_headers(), headers=self._OCS_HEADERS_JSON,
) )
data = self._unwrap_ocs(response.json()) data = self._unwrap_ocs(response.json())
return data["page"] return data["page"]
@@ -204,7 +209,7 @@ class CollectivesClient(BaseNextcloudClient):
"GET", "GET",
f"{API_BASE}/collectives/{collective_id}/search", f"{API_BASE}/collectives/{collective_id}/search",
params={"searchString": query}, params={"searchString": query},
headers=self._get_ocs_headers(), headers=self._OCS_HEADERS,
) )
data = self._unwrap_ocs(response.json()) data = self._unwrap_ocs(response.json())
return data["pages"] return data["pages"]
@@ -216,7 +221,7 @@ class CollectivesClient(BaseNextcloudClient):
response = await self._make_request( response = await self._make_request(
"GET", "GET",
f"{API_BASE}/collectives/{collective_id}/tags", f"{API_BASE}/collectives/{collective_id}/tags",
headers=self._get_ocs_headers(), headers=self._OCS_HEADERS,
) )
data = self._unwrap_ocs(response.json()) data = self._unwrap_ocs(response.json())
return data["tags"] return data["tags"]
@@ -230,7 +235,7 @@ class CollectivesClient(BaseNextcloudClient):
"POST", "POST",
f"{API_BASE}/collectives/{collective_id}/tags", f"{API_BASE}/collectives/{collective_id}/tags",
json=json_data, json=json_data,
headers=self._get_ocs_headers(), headers=self._OCS_HEADERS_JSON,
) )
data = self._unwrap_ocs(response.json()) data = self._unwrap_ocs(response.json())
return data["tag"] return data["tag"]
@@ -240,7 +245,7 @@ class CollectivesClient(BaseNextcloudClient):
response = await self._make_request( response = await self._make_request(
"PUT", "PUT",
f"{API_BASE}/collectives/{collective_id}/pages/{page_id}/tags/{tag_id}", f"{API_BASE}/collectives/{collective_id}/pages/{page_id}/tags/{tag_id}",
headers=self._get_ocs_headers(), headers=self._OCS_HEADERS_JSON,
) )
self._unwrap_ocs(response.json()) self._unwrap_ocs(response.json())
@@ -249,7 +254,7 @@ class CollectivesClient(BaseNextcloudClient):
response = await self._make_request( response = await self._make_request(
"DELETE", "DELETE",
f"{API_BASE}/collectives/{collective_id}/pages/{page_id}/tags/{tag_id}", f"{API_BASE}/collectives/{collective_id}/pages/{page_id}/tags/{tag_id}",
headers=self._get_ocs_headers(), headers=self._OCS_HEADERS,
) )
self._unwrap_ocs(response.json()) self._unwrap_ocs(response.json())
@@ -260,7 +265,7 @@ class CollectivesClient(BaseNextcloudClient):
response = await self._make_request( response = await self._make_request(
"GET", "GET",
f"{API_BASE}/collectives/{collective_id}/pages/trash", f"{API_BASE}/collectives/{collective_id}/pages/trash",
headers=self._get_ocs_headers(), headers=self._OCS_HEADERS,
) )
data = self._unwrap_ocs(response.json()) data = self._unwrap_ocs(response.json())
return data["pages"] return data["pages"]
@@ -270,7 +275,7 @@ class CollectivesClient(BaseNextcloudClient):
response = await self._make_request( response = await self._make_request(
"PATCH", "PATCH",
f"{API_BASE}/collectives/{collective_id}/pages/trash/{page_id}", f"{API_BASE}/collectives/{collective_id}/pages/trash/{page_id}",
headers=self._get_ocs_headers(), headers=self._OCS_HEADERS,
) )
data = self._unwrap_ocs(response.json()) data = self._unwrap_ocs(response.json())
return data["page"] return data["page"]
+1 -1
View File
@@ -249,7 +249,7 @@ def configure_collectives_tools(mcp: FastMCP):
Args: Args:
collective_id: ID of the collective collective_id: ID of the collective
emoji: New emoji for the collective (required) emoji: New emoji for the collective
""" """
client = await get_client(ctx) client = await get_client(ctx)
try: try:
@@ -117,7 +117,7 @@ async def test_create_collective(mocker):
async def test_trash_collective(mocker): async def test_trash_collective(mocker):
"""Test trashing a collective sends DELETE to correct endpoint.""" """Test trashing a collective sends DELETE to correct endpoint."""
mock_response = create_mock_response(status_code=200, json_data={}) mock_response = _ocs_response({})
mock_request = mocker.patch.object( mock_request = mocker.patch.object(
CollectivesClient, "_make_request", return_value=mock_response CollectivesClient, "_make_request", return_value=mock_response
) )
@@ -133,7 +133,7 @@ async def test_trash_collective(mocker):
async def test_delete_collective(mocker): async def test_delete_collective(mocker):
"""Test permanently deleting a collective sends DELETE to trash endpoint.""" """Test permanently deleting a collective sends DELETE to trash endpoint."""
mock_response = create_mock_response(status_code=200, json_data={}) mock_response = _ocs_response({})
mock_request = mocker.patch.object( mock_request = mocker.patch.object(
CollectivesClient, "_make_request", return_value=mock_response CollectivesClient, "_make_request", return_value=mock_response
) )
@@ -203,7 +203,7 @@ async def test_create_page(mocker):
async def test_trash_page(mocker): async def test_trash_page(mocker):
"""Test trashing a page sends DELETE.""" """Test trashing a page sends DELETE."""
mock_response = create_mock_response(status_code=200, json_data={}) mock_response = _ocs_response({})
mock_request = mocker.patch.object( mock_request = mocker.patch.object(
CollectivesClient, "_make_request", return_value=mock_response CollectivesClient, "_make_request", return_value=mock_response
) )
@@ -358,8 +358,8 @@ async def test_restore_page(mocker):
# --- Error Handling --- # --- Error Handling ---
async def test_ocs_missing_data_returns_empty(mocker): async def test_ocs_missing_data_raises_key_error(mocker):
"""Test that OCS envelope without 'data' key returns empty dict.""" """Test that OCS envelope without 'data' key causes KeyError on field access."""
mock_response = create_mock_response( mock_response = create_mock_response(
status_code=200, status_code=200,
json_data={ json_data={
@@ -371,12 +371,25 @@ async def test_ocs_missing_data_returns_empty(mocker):
mocker.patch.object(CollectivesClient, "_make_request", return_value=mock_response) mocker.patch.object(CollectivesClient, "_make_request", return_value=mock_response)
client = CollectivesClient(mocker.AsyncMock(spec=httpx.AsyncClient), "testuser") client = CollectivesClient(mocker.AsyncMock(spec=httpx.AsyncClient), "testuser")
# get_collectives accesses data["collectives"], which will KeyError on empty dict # _unwrap_ocs returns {} when "data" is absent; the caller then
# This tests that _unwrap_ocs itself doesn't crash — it returns {} # raises KeyError when accessing the expected key (e.g. "collectives")
with pytest.raises(KeyError): with pytest.raises(KeyError):
await client.get_collectives() await client.get_collectives()
async def test_non_ocs_envelope_raises_ocs_error(mocker):
"""Test that a non-OCS response (e.g. proxy error) raises OCSError."""
mock_response = create_mock_response(
status_code=200,
json_data={"error": "Bad Gateway"},
)
mocker.patch.object(CollectivesClient, "_make_request", return_value=mock_response)
client = CollectivesClient(mocker.AsyncMock(spec=httpx.AsyncClient), "testuser")
with pytest.raises(OCSError, match="not an OCS envelope"):
await client.get_collectives()
async def test_ocs_error_status_raises(mocker): async def test_ocs_error_status_raises(mocker):
"""Test that OCS envelope with error statuscode raises OCSError.""" """Test that OCS envelope with error statuscode raises OCSError."""
mock_response = create_mock_response( mock_response = create_mock_response(