fix: encode dynamic URL path segments (#5963)
Co-authored-by: kartik-mem0 <kartik.labhshetwar@mem0.ai>
This commit is contained in:
@@ -3,6 +3,7 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
from urllib.parse import quote
|
||||
|
||||
import httpx
|
||||
|
||||
@@ -11,6 +12,10 @@ from mem0_cli.backend.base import Backend
|
||||
from mem0_cli.config import PlatformConfig
|
||||
|
||||
|
||||
def _encode_path_segment(value: Any) -> str:
|
||||
return quote(str(value), safe="")
|
||||
|
||||
|
||||
class PlatformBackend(Backend):
|
||||
"""Backend that talks to the mem0 Platform API."""
|
||||
|
||||
@@ -196,7 +201,11 @@ class PlatformBackend(Backend):
|
||||
)
|
||||
|
||||
def get(self, memory_id: str) -> dict:
|
||||
return self._request("GET", f"/v1/memories/{memory_id}/", params={"source": "CLI"})
|
||||
return self._request(
|
||||
"GET",
|
||||
f"/v1/memories/{_encode_path_segment(memory_id)}/",
|
||||
params={"source": "CLI"},
|
||||
)
|
||||
|
||||
def list_memories(
|
||||
self,
|
||||
@@ -250,7 +259,11 @@ class PlatformBackend(Backend):
|
||||
if metadata:
|
||||
payload["metadata"] = metadata
|
||||
payload["source"] = "CLI"
|
||||
return self._request("PUT", f"/v1/memories/{memory_id}/", json=payload)
|
||||
return self._request(
|
||||
"PUT",
|
||||
f"/v1/memories/{_encode_path_segment(memory_id)}/",
|
||||
json=payload,
|
||||
)
|
||||
|
||||
def delete(
|
||||
self,
|
||||
@@ -274,7 +287,11 @@ class PlatformBackend(Backend):
|
||||
params["run_id"] = run_id
|
||||
return self._request("DELETE", "/v1/memories/", params=params)
|
||||
elif memory_id:
|
||||
return self._request("DELETE", f"/v1/memories/{memory_id}/", params={"source": "CLI"})
|
||||
return self._request(
|
||||
"DELETE",
|
||||
f"/v1/memories/{_encode_path_segment(memory_id)}/",
|
||||
params={"source": "CLI"},
|
||||
)
|
||||
else:
|
||||
raise ValueError("Either memory_id or --all is required")
|
||||
|
||||
@@ -302,7 +319,9 @@ class PlatformBackend(Backend):
|
||||
results: dict = {}
|
||||
for entity_type, entity_id in entities.items():
|
||||
results[entity_type] = self._request(
|
||||
"DELETE", f"/v2/entities/{entity_type}/{entity_id}/", params={"source": "CLI"}
|
||||
"DELETE",
|
||||
f"/v2/entities/{_encode_path_segment(entity_type)}/{_encode_path_segment(entity_id)}/",
|
||||
params={"source": "CLI"},
|
||||
)
|
||||
return results
|
||||
|
||||
@@ -348,7 +367,7 @@ class PlatformBackend(Backend):
|
||||
return result if isinstance(result, list) else result.get("results", [])
|
||||
|
||||
def get_event(self, event_id: str) -> dict:
|
||||
return self._request("GET", f"/v1/event/{event_id}/")
|
||||
return self._request("GET", f"/v1/event/{_encode_path_segment(event_id)}/")
|
||||
|
||||
|
||||
class AuthError(Exception):
|
||||
|
||||
@@ -0,0 +1,43 @@
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
from mem0_cli.backend.platform import PlatformBackend
|
||||
|
||||
|
||||
def _backend(sample_config):
|
||||
backend = PlatformBackend(sample_config.platform)
|
||||
backend._client = MagicMock()
|
||||
backend._client.request.return_value = MagicMock(
|
||||
status_code=200,
|
||||
json=lambda: {"message": "ok"},
|
||||
headers={},
|
||||
raise_for_status=lambda: None,
|
||||
)
|
||||
return backend
|
||||
|
||||
|
||||
def test_memory_id_path_segments_are_encoded(sample_config):
|
||||
backend = _backend(sample_config)
|
||||
|
||||
backend.get("mem/a?b#c")
|
||||
backend.update("mem/a?b#c", content="updated")
|
||||
backend.delete("mem/a?b#c")
|
||||
|
||||
paths = [call.args[1] for call in backend._client.request.call_args_list]
|
||||
assert paths == [
|
||||
"/v1/memories/mem%2Fa%3Fb%23c/",
|
||||
"/v1/memories/mem%2Fa%3Fb%23c/",
|
||||
"/v1/memories/mem%2Fa%3Fb%23c/",
|
||||
]
|
||||
|
||||
|
||||
def test_entity_and_event_path_segments_are_encoded(sample_config):
|
||||
backend = _backend(sample_config)
|
||||
|
||||
backend.delete_entities(user_id="org/team?active#frag")
|
||||
backend.get_event("evt/a?b#c")
|
||||
|
||||
paths = [call.args[1] for call in backend._client.request.call_args_list]
|
||||
assert paths == [
|
||||
"/v2/entities/user/org%2Fteam%3Factive%23frag/",
|
||||
"/v1/event/evt%2Fa%3Fb%23c/",
|
||||
]
|
||||
Reference in New Issue
Block a user