681 lines
25 KiB
Python
681 lines
25 KiB
Python
"""Tests for Mem0Backend abstraction — PlatformBackend, OSSBackend, SelfHostedBackend."""
|
|
|
|
import copy
|
|
import importlib
|
|
import json
|
|
import os
|
|
import sys
|
|
import types
|
|
from dataclasses import dataclass, field
|
|
from types import SimpleNamespace
|
|
|
|
import pytest
|
|
|
|
from plugins.memory.mem0._backend import (
|
|
Mem0Backend,
|
|
PlatformBackend,
|
|
OSSBackend,
|
|
SelfHostedBackend,
|
|
)
|
|
|
|
|
|
class FakePlatformClient:
|
|
"""Fake MemoryClient for PlatformBackend tests."""
|
|
|
|
def __init__(self):
|
|
self.calls = []
|
|
|
|
def search(self, query, **kwargs):
|
|
self.calls.append(("search", query, kwargs))
|
|
return {"results": [{"id": "m1", "memory": "fact1", "score": 0.9}]}
|
|
|
|
def get_all(self, **kwargs):
|
|
self.calls.append(("get_all", kwargs))
|
|
return {"count": 1, "next": None, "results": [{"id": "m1", "memory": "fact1"}]}
|
|
|
|
def add(self, messages, **kwargs):
|
|
self.calls.append(("add", messages, kwargs))
|
|
return {"status": "PENDING", "event_id": "evt-1"}
|
|
|
|
def update(self, **kwargs):
|
|
self.calls.append(("update", kwargs))
|
|
return {"id": kwargs["memory_id"], "text": kwargs["text"]}
|
|
|
|
def delete(self, **kwargs):
|
|
self.calls.append(("delete", kwargs))
|
|
|
|
|
|
class TestPlatformBackend:
|
|
|
|
def _make(self):
|
|
client = FakePlatformClient()
|
|
backend = PlatformBackend.__new__(PlatformBackend)
|
|
backend._client = client
|
|
return backend, client
|
|
|
|
def test_search_forwards_params(self):
|
|
backend, client = self._make()
|
|
result = backend.search("test query", filters={"user_id": "u1"}, top_k=5)
|
|
assert client.calls[0][0] == "search"
|
|
assert client.calls[0][1] == "test query"
|
|
assert client.calls[0][2]["filters"] == {"user_id": "u1"}
|
|
assert client.calls[0][2]["top_k"] == 5
|
|
|
|
|
|
def test_add_forwards_kwargs(self):
|
|
backend, client = self._make()
|
|
msgs = [{"role": "user", "content": "hi"}]
|
|
result = backend.add(msgs, user_id="u1", agent_id="hermes", infer=False)
|
|
call = client.calls[0]
|
|
assert call[2]["user_id"] == "u1"
|
|
assert call[2]["infer"] is False
|
|
# metadata kwarg should be omitted entirely when not provided so we
|
|
# don't surprise older mem0 client versions with an unknown kwarg.
|
|
assert "metadata" not in call[2]
|
|
|
|
|
|
def test_update_forwards(self):
|
|
backend, client = self._make()
|
|
backend.update("m1", "new text")
|
|
assert client.calls[0][1] == {"memory_id": "m1", "text": "new text"}
|
|
|
|
def test_delete_forwards(self):
|
|
backend, client = self._make()
|
|
backend.delete("m1")
|
|
assert client.calls[0][1] == {"memory_id": "m1"}
|
|
|
|
|
|
class FakeOSSMemory:
|
|
"""Fake mem0.Memory for OSSBackend tests."""
|
|
|
|
def __init__(self):
|
|
self.calls = []
|
|
|
|
def search(self, query, **kwargs):
|
|
self.calls.append(("search", query, kwargs))
|
|
return {"results": [{"id": "m1", "memory": "fact1", "score": 0.8}]}
|
|
|
|
def get_all(self, **kwargs):
|
|
self.calls.append(("get_all", kwargs))
|
|
return {"results": [{"id": "m1", "memory": "fact1"}]}
|
|
|
|
def add(self, messages, **kwargs):
|
|
self.calls.append(("add", messages, kwargs))
|
|
return {"results": [{"id": "m1", "memory": "fact1", "event": "ADD"}]}
|
|
|
|
def update(self, memory_id, **kwargs):
|
|
self.calls.append(("update", memory_id, kwargs))
|
|
return {"message": "Memory updated successfully!"}
|
|
|
|
def delete(self, memory_id):
|
|
self.calls.append(("delete", memory_id))
|
|
return {"message": "Memory deleted successfully!"}
|
|
|
|
|
|
@dataclass
|
|
class _FakeMem0State:
|
|
factory_registrations: list = field(default_factory=list)
|
|
from_config_calls: int = 0
|
|
clients: list = field(default_factory=list)
|
|
requests: list = field(default_factory=list)
|
|
|
|
|
|
def _install_fake_mem0(monkeypatch):
|
|
"""Install a small mem0 2.0.10-shaped surface for OSS backend tests."""
|
|
|
|
state = _FakeMem0State()
|
|
|
|
class BaseLlmConfig:
|
|
def __init__(
|
|
self,
|
|
model=None,
|
|
temperature=0.1,
|
|
api_key=None,
|
|
max_tokens=2000,
|
|
top_p=0.1,
|
|
top_k=1,
|
|
enable_vision=False,
|
|
vision_details="auto",
|
|
reasoning_effort=None,
|
|
http_client_proxies=None,
|
|
is_reasoning_model=None,
|
|
**kwargs,
|
|
):
|
|
self.model = model
|
|
self.temperature = temperature
|
|
self.api_key = api_key
|
|
self.max_tokens = max_tokens
|
|
self.top_p = top_p
|
|
self.top_k = top_k
|
|
self.enable_vision = enable_vision
|
|
self.vision_details = vision_details
|
|
self.reasoning_effort = reasoning_effort
|
|
self.http_client_proxies = http_client_proxies
|
|
self.is_reasoning_model = is_reasoning_model
|
|
for name, value in kwargs.items():
|
|
setattr(self, name, value)
|
|
|
|
class OpenAIConfig(BaseLlmConfig):
|
|
def __init__(
|
|
self,
|
|
model=None,
|
|
temperature=0.1,
|
|
api_key=None,
|
|
max_tokens=2000,
|
|
top_p=0.1,
|
|
top_k=1,
|
|
enable_vision=False,
|
|
vision_details="auto",
|
|
reasoning_effort=None,
|
|
http_client_proxies=None,
|
|
is_reasoning_model=None,
|
|
openai_base_url=None,
|
|
models=None,
|
|
route="fallback",
|
|
openrouter_base_url=None,
|
|
site_url=None,
|
|
app_name=None,
|
|
store=None,
|
|
response_callback=None,
|
|
):
|
|
super().__init__(
|
|
model=model,
|
|
temperature=temperature,
|
|
api_key=api_key,
|
|
max_tokens=max_tokens,
|
|
top_p=top_p,
|
|
top_k=top_k,
|
|
enable_vision=enable_vision,
|
|
vision_details=vision_details,
|
|
reasoning_effort=reasoning_effort,
|
|
http_client_proxies=http_client_proxies,
|
|
is_reasoning_model=is_reasoning_model,
|
|
)
|
|
self.openai_base_url = openai_base_url
|
|
self.models = models
|
|
self.route = route
|
|
self.openrouter_base_url = openrouter_base_url
|
|
self.site_url = site_url
|
|
self.app_name = app_name
|
|
self.store = store
|
|
self.response_callback = response_callback
|
|
|
|
class LLMBase:
|
|
def __init__(self, config=None):
|
|
self.config = config or BaseLlmConfig()
|
|
if not hasattr(self.config, "model"):
|
|
raise ValueError("Configuration must have a 'model' attribute")
|
|
|
|
def _get_supported_params(self, **kwargs):
|
|
if self.config.is_reasoning_model:
|
|
return {
|
|
name: kwargs[name]
|
|
for name in ("messages", "response_format", "tools", "tool_choice")
|
|
if name in kwargs
|
|
}
|
|
params = {
|
|
"temperature": self.config.temperature,
|
|
"top_p": self.config.top_p,
|
|
"max_tokens": self.config.max_tokens,
|
|
}
|
|
params.update(kwargs)
|
|
return params
|
|
|
|
class OpenAILLM(LLMBase):
|
|
@staticmethod
|
|
def _parse_response(response, tools):
|
|
if not tools:
|
|
return response.choices[0].message.content
|
|
parsed = {
|
|
"content": response.choices[0].message.content,
|
|
"tool_calls": [],
|
|
}
|
|
for tool_call in response.choices[0].message.tool_calls or []:
|
|
parsed["tool_calls"].append(
|
|
{
|
|
"name": tool_call.function.name,
|
|
"arguments": json.loads(tool_call.function.arguments),
|
|
}
|
|
)
|
|
return parsed
|
|
|
|
class Factory:
|
|
provider_to_class = {
|
|
"openai": ("mem0.llms.openai.OpenAILLM", OpenAIConfig),
|
|
"ollama": ("mem0.llms.openai.OpenAILLM", BaseLlmConfig),
|
|
}
|
|
|
|
@classmethod
|
|
def register_provider(cls, name, class_path, config_class=None):
|
|
cls.provider_to_class[name] = (
|
|
class_path,
|
|
config_class or BaseLlmConfig,
|
|
)
|
|
state.factory_registrations.append((name, class_path, config_class))
|
|
|
|
@classmethod
|
|
def create(cls, provider_name, config=None, **kwargs):
|
|
class_path, config_class = cls.provider_to_class[provider_name]
|
|
if config is None:
|
|
config = config_class(**kwargs)
|
|
elif isinstance(config, dict):
|
|
config = config_class(**config)
|
|
module_name, class_name = class_path.rsplit(".", 1)
|
|
llm_class = getattr(importlib.import_module(module_name), class_name)
|
|
return llm_class(config)
|
|
|
|
class MemoryConfig:
|
|
def __init__(self, **config):
|
|
llm = config["llm"]
|
|
if llm["provider"] not in {"openai", "ollama"}:
|
|
raise ValueError(
|
|
f"Unsupported LLM provider: {llm['provider']}"
|
|
)
|
|
self.llm = SimpleNamespace(
|
|
provider=llm["provider"],
|
|
config=copy.deepcopy(llm.get("config", {})),
|
|
)
|
|
embedder = config["embedder"]
|
|
self.embedder = SimpleNamespace(
|
|
provider=embedder["provider"],
|
|
config=copy.deepcopy(embedder.get("config", {})),
|
|
)
|
|
vector_store = config["vector_store"]
|
|
self.vector_store = SimpleNamespace(
|
|
provider=vector_store["provider"],
|
|
config=copy.deepcopy(vector_store.get("config", {})),
|
|
)
|
|
self.version = config.get("version", "v1.1")
|
|
|
|
class Memory:
|
|
instances = []
|
|
|
|
def __init__(self, config):
|
|
self.config = config
|
|
self.llm = Factory.create(config.llm.provider, config.llm.config)
|
|
self.embedding_model = SimpleNamespace(
|
|
provider=config.embedder.provider,
|
|
config=config.embedder.config,
|
|
)
|
|
self.vector_store = SimpleNamespace(
|
|
provider=config.vector_store.provider,
|
|
config=config.vector_store.config,
|
|
)
|
|
type(self).instances.append(self)
|
|
|
|
@classmethod
|
|
def from_config(cls, config):
|
|
# This mirrors mem0 2.0.10: validation rejects the private provider
|
|
# before the factory gets a chance to resolve its registration.
|
|
state.from_config_calls += 1
|
|
return cls(MemoryConfig(**config))
|
|
|
|
class FakeOpenAI:
|
|
def __init__(self, *, api_key, base_url):
|
|
self.api_key = api_key
|
|
self.base_url = base_url
|
|
state.clients.append(self)
|
|
self.chat = SimpleNamespace(
|
|
completions=SimpleNamespace(create=self._create)
|
|
)
|
|
|
|
def _create(self, **params):
|
|
state.requests.append(params)
|
|
return SimpleNamespace(
|
|
choices=[
|
|
SimpleNamespace(
|
|
message=SimpleNamespace(
|
|
content="direct answer",
|
|
tool_calls=[
|
|
SimpleNamespace(
|
|
function=SimpleNamespace(
|
|
name="remember",
|
|
arguments='{"fact": "tea"}',
|
|
)
|
|
)
|
|
],
|
|
)
|
|
)
|
|
]
|
|
)
|
|
|
|
package_names = {
|
|
"mem0": types.ModuleType("mem0"),
|
|
"mem0.configs": types.ModuleType("mem0.configs"),
|
|
"mem0.configs.llms": types.ModuleType("mem0.configs.llms"),
|
|
"mem0.llms": types.ModuleType("mem0.llms"),
|
|
"mem0.utils": types.ModuleType("mem0.utils"),
|
|
"mem0.configs.base": types.ModuleType("mem0.configs.base"),
|
|
"mem0.configs.llms.base": types.ModuleType("mem0.configs.llms.base"),
|
|
"mem0.configs.llms.openai": types.ModuleType("mem0.configs.llms.openai"),
|
|
"mem0.llms.base": types.ModuleType("mem0.llms.base"),
|
|
"mem0.llms.openai": types.ModuleType("mem0.llms.openai"),
|
|
"mem0.utils.factory": types.ModuleType("mem0.utils.factory"),
|
|
"openai": types.ModuleType("openai"),
|
|
}
|
|
setattr(package_names["mem0"], "Memory", Memory)
|
|
setattr(package_names["mem0.configs.base"], "MemoryConfig", MemoryConfig)
|
|
setattr(package_names["mem0.configs.llms.base"], "BaseLlmConfig", BaseLlmConfig)
|
|
setattr(package_names["mem0.configs.llms.openai"], "OpenAIConfig", OpenAIConfig)
|
|
setattr(package_names["mem0.llms.base"], "LLMBase", LLMBase)
|
|
setattr(package_names["mem0.llms.openai"], "OpenAILLM", OpenAILLM)
|
|
setattr(package_names["mem0.utils.factory"], "LlmFactory", Factory)
|
|
setattr(package_names["openai"], "OpenAI", FakeOpenAI)
|
|
for name, module in package_names.items():
|
|
if name in {"mem0", "mem0.configs", "mem0.configs.llms", "mem0.llms", "mem0.utils"}:
|
|
module.__path__ = []
|
|
monkeypatch.setitem(sys.modules, name, module)
|
|
|
|
# The class-path registration imports this module after the fake mem0
|
|
# surface is installed, so it binds to the test doubles above.
|
|
monkeypatch.delitem(
|
|
sys.modules, "plugins.memory.mem0._openai_llm", raising=False
|
|
)
|
|
return state, Memory, Factory
|
|
|
|
|
|
class TestOSSBackend:
|
|
|
|
def _make(self):
|
|
memory = FakeOSSMemory()
|
|
backend = OSSBackend.__new__(OSSBackend)
|
|
backend._memory = memory
|
|
return backend, memory
|
|
|
|
|
|
def test_legacy_api_base_aliases_are_normalized_before_mem0_init(self, monkeypatch):
|
|
state, Memory, factory = _install_fake_mem0(monkeypatch)
|
|
raw = {
|
|
"llm": {
|
|
"provider": "openai",
|
|
"config": {
|
|
"model": "gpt-5-mini",
|
|
"api_key": "openai-sentinel",
|
|
"api_base": "https://llm.example/v1",
|
|
},
|
|
},
|
|
"embedder": {
|
|
"provider": "ollama",
|
|
"config": {"model": "nomic-embed-text", "api_base": "http://ollama:11434"},
|
|
},
|
|
"vector_store": {"provider": "qdrant", "config": {}},
|
|
}
|
|
before = copy.deepcopy(raw)
|
|
environment = dict(os.environ)
|
|
|
|
OSSBackend(raw)
|
|
|
|
assert len(Memory.instances) == 1
|
|
captured = Memory.instances[0].config
|
|
assert captured.llm.provider == "hermes_openai"
|
|
assert captured.llm.config["openai_base_url"] == "https://llm.example/v1"
|
|
assert captured.embedder.provider == "ollama"
|
|
assert captured.embedder.config["ollama_base_url"] == "http://ollama:11434"
|
|
assert "api_base" not in captured.llm.config
|
|
assert "api_base" not in captured.embedder.config
|
|
assert factory.provider_to_class["hermes_openai"][1].__name__ == "OpenAIConfig"
|
|
assert len(state.factory_registrations) == 1
|
|
assert state.from_config_calls == 0
|
|
assert raw == before
|
|
assert dict(os.environ) == environment
|
|
|
|
def test_direct_openai_uses_openai_credentials_and_request_shape(self, monkeypatch):
|
|
state, _, factory = _install_fake_mem0(monkeypatch)
|
|
monkeypatch.setenv("OPENROUTER_API_KEY", "router-sentinel")
|
|
monkeypatch.setenv("OPENAI_API_KEY", "env-openai-sentinel")
|
|
|
|
module = importlib.import_module("plugins.memory.mem0._openai_llm")
|
|
callback_calls = []
|
|
config = factory.provider_to_class["openai"][1](
|
|
model="gpt-5-mini",
|
|
api_key="configured-openai-sentinel",
|
|
openai_base_url="https://openai.example/v1",
|
|
models=["router-model"],
|
|
route="lowest-latency",
|
|
site_url="https://hermes.example",
|
|
app_name="Hermes",
|
|
store=True,
|
|
response_callback=lambda *args: callback_calls.append(args),
|
|
)
|
|
adapter = module.DirectOpenAILLM(config)
|
|
assert adapter.config.is_reasoning_model is True
|
|
tools = [
|
|
{
|
|
"type": "function",
|
|
"function": {"name": "remember", "parameters": {}},
|
|
}
|
|
]
|
|
|
|
result = adapter.generate_response(
|
|
[{"role": "user", "content": "remember tea"}],
|
|
response_format={"type": "json_object"},
|
|
tools=tools,
|
|
tool_choice="required",
|
|
)
|
|
|
|
assert len(state.clients) == 1
|
|
client = state.clients[0]
|
|
assert client.api_key == "configured-openai-sentinel"
|
|
assert client.base_url == "https://openai.example/v1"
|
|
request = state.requests[0]
|
|
assert request["model"] == "gpt-5-mini"
|
|
assert request["tools"] == tools
|
|
assert request["tool_choice"] == "required"
|
|
assert request["response_format"] == {"type": "json_object"}
|
|
assert request["store"] is True
|
|
assert "models" not in request
|
|
assert "route" not in request
|
|
assert "extra_headers" not in request
|
|
assert "temperature" not in request
|
|
assert "top_p" not in request
|
|
assert "max_tokens" not in request
|
|
assert result == {
|
|
"content": "direct answer",
|
|
"tool_calls": [{"name": "remember", "arguments": {"fact": "tea"}}],
|
|
}
|
|
assert len(callback_calls) == 1
|
|
assert callback_calls[0][0] is adapter
|
|
assert callback_calls[0][2] == request
|
|
|
|
def test_direct_openai_preserves_explicit_non_reasoning_override(self, monkeypatch):
|
|
state, _, factory = _install_fake_mem0(monkeypatch)
|
|
config = factory.provider_to_class["openai"][1](
|
|
model="gpt-5-mini",
|
|
api_key="configured-openai-sentinel",
|
|
is_reasoning_model=False,
|
|
)
|
|
|
|
module = importlib.import_module("plugins.memory.mem0._openai_llm")
|
|
adapter = module.DirectOpenAILLM(config)
|
|
adapter.generate_response([{"role": "user", "content": "remember tea"}])
|
|
|
|
assert adapter.config.is_reasoning_model is False
|
|
request = state.requests[0]
|
|
assert request["temperature"] == 0.1
|
|
assert request["top_p"] == 0.1
|
|
assert request["max_tokens"] == 2000
|
|
|
|
def test_direct_openai_defaults_missing_model_to_reasoning_safe_mini(self, monkeypatch):
|
|
monkeypatch.setenv("OPENAI_API_KEY", "environment-openai-sentinel")
|
|
_install_fake_mem0(monkeypatch)
|
|
|
|
module = importlib.import_module("plugins.memory.mem0._openai_llm")
|
|
adapter = module.DirectOpenAILLM()
|
|
|
|
assert adapter.config.model == "gpt-5-mini"
|
|
assert adapter.config.is_reasoning_model is True
|
|
|
|
def test_direct_openai_uses_openai_environment_when_config_omits_values(self, monkeypatch):
|
|
state, _, factory = _install_fake_mem0(monkeypatch)
|
|
monkeypatch.setenv("OPENROUTER_API_KEY", "router-sentinel")
|
|
monkeypatch.setenv("OPENAI_API_KEY", "env-openai-sentinel")
|
|
monkeypatch.setenv("OPENAI_BASE_URL", "https://env-openai.example/v1")
|
|
|
|
module = importlib.import_module("plugins.memory.mem0._openai_llm")
|
|
config = factory.provider_to_class["openai"][1](model="gpt-5-mini")
|
|
adapter = module.DirectOpenAILLM(config)
|
|
|
|
assert len(state.clients) == 1
|
|
assert state.clients[0].api_key == "env-openai-sentinel"
|
|
assert state.clients[0].base_url == "https://env-openai.example/v1"
|
|
|
|
def test_missing_openai_key_fails_before_client_and_hides_router_secret(self, monkeypatch):
|
|
state, _, factory = _install_fake_mem0(monkeypatch)
|
|
router_secret = "router-secret-sentinel"
|
|
monkeypatch.setenv("OPENROUTER_API_KEY", router_secret)
|
|
monkeypatch.delenv("OPENAI_API_KEY", raising=False)
|
|
|
|
module = importlib.import_module("plugins.memory.mem0._openai_llm")
|
|
config = factory.provider_to_class["openai"][1](
|
|
model="gpt-5-mini",
|
|
api_key=None,
|
|
)
|
|
|
|
with pytest.raises(ValueError) as exc_info:
|
|
module.DirectOpenAILLM(config)
|
|
|
|
assert "OpenAI API key" in str(exc_info.value)
|
|
assert router_secret not in str(exc_info.value)
|
|
assert state.clients == []
|
|
assert state.requests == []
|
|
|
|
def test_registration_is_idempotent_and_clients_keep_instance_config(self, monkeypatch):
|
|
state, Memory, factory = _install_fake_mem0(monkeypatch)
|
|
first = {
|
|
"llm": {
|
|
"provider": "openai",
|
|
"config": {
|
|
"model": "gpt-5-mini",
|
|
"api_key": "first-openai-sentinel",
|
|
"openai_base_url": "https://first.example/v1",
|
|
},
|
|
},
|
|
"embedder": {"provider": "ollama", "config": {}},
|
|
"vector_store": {"provider": "qdrant", "config": {}},
|
|
}
|
|
second = {
|
|
"llm": {
|
|
"provider": "openai",
|
|
"config": {
|
|
"model": "gpt-5-mini",
|
|
"api_key": "second-openai-sentinel",
|
|
"openai_base_url": "https://second.example/v1",
|
|
},
|
|
},
|
|
"embedder": {"provider": "ollama", "config": {}},
|
|
"vector_store": {"provider": "qdrant", "config": {}},
|
|
}
|
|
first_before = copy.deepcopy(first)
|
|
second_before = copy.deepcopy(second)
|
|
|
|
OSSBackend(first)
|
|
OSSBackend(second)
|
|
|
|
assert len(state.factory_registrations) == 1
|
|
assert factory.provider_to_class["hermes_openai"][0].endswith(
|
|
"_openai_llm.DirectOpenAILLM"
|
|
)
|
|
assert [
|
|
(client.api_key, client.base_url) for client in state.clients
|
|
] == [
|
|
("first-openai-sentinel", "https://first.example/v1"),
|
|
("second-openai-sentinel", "https://second.example/v1"),
|
|
]
|
|
assert len(Memory.instances) == 2
|
|
assert state.from_config_calls == 0
|
|
assert first == first_before
|
|
assert second == second_before
|
|
|
|
def test_ollama_bypasses_direct_openai_adapter(self, monkeypatch):
|
|
state, Memory, factory = _install_fake_mem0(monkeypatch)
|
|
raw = {
|
|
"llm": {
|
|
"provider": "ollama",
|
|
"config": {
|
|
"model": "llama3.1:8b",
|
|
"api_base": "http://ollama:11434",
|
|
},
|
|
},
|
|
"embedder": {
|
|
"provider": "ollama",
|
|
"config": {
|
|
"model": "nomic-embed-text",
|
|
"api_base": "http://ollama:11434",
|
|
},
|
|
},
|
|
"vector_store": {"provider": "qdrant", "config": {}},
|
|
}
|
|
before = copy.deepcopy(raw)
|
|
|
|
OSSBackend(raw)
|
|
|
|
assert len(Memory.instances) == 1
|
|
assert state.from_config_calls == 1
|
|
assert Memory.instances[0].config.llm.provider == "ollama"
|
|
assert Memory.instances[0].config.embedder.provider == "ollama"
|
|
assert "hermes_openai" not in factory.provider_to_class
|
|
assert state.clients == []
|
|
assert raw == before
|
|
|
|
|
|
httpx = pytest.importorskip("httpx")
|
|
|
|
|
|
class _StubServer:
|
|
"""Records requests and serves the real self-hosted server's response shapes."""
|
|
|
|
def __init__(self, rows=10):
|
|
self.requests = []
|
|
self._rows = [{"id": f"m{i}", "memory": f"f{i}"} for i in range(rows)]
|
|
|
|
def handler(self, request):
|
|
self.requests.append(request)
|
|
path, method = request.url.path, request.method
|
|
if path == "/search" and method == "POST":
|
|
return httpx.Response(200, json={"results": [{"id": "m1", "memory": "tea", "score": 0.9}]})
|
|
if path == "/memories" and method == "GET":
|
|
top_k = int(request.url.params.get("top_k", len(self._rows)))
|
|
return httpx.Response(200, json={"results": self._rows[:top_k]})
|
|
if path == "/memories" and method == "POST":
|
|
return httpx.Response(200, json={"results": [{"id": "new", "memory": "stored", "event": "ADD"}]})
|
|
if path.startswith("/memories/") and method in ("PUT", "DELETE"):
|
|
if path.endswith("/missing"): # server 404s unknown ids
|
|
return httpx.Response(404, json={"detail": "Memory not found"})
|
|
verb = "updated" if method == "PUT" else "Memory deleted successfully"
|
|
return httpx.Response(200, json={"message": verb})
|
|
return httpx.Response(404, json={"detail": "not found"})
|
|
|
|
|
|
def _backend(server, api_key="adminkey", host="http://sh:8888"):
|
|
"""Build a SelfHostedBackend routed through the stub transport.
|
|
|
|
Uses the real __init__ (via the injectable ``transport`` kwarg) so the
|
|
constructor's header/base_url setup is exercised by every test here.
|
|
"""
|
|
return SelfHostedBackend(
|
|
api_key, host, transport=httpx.MockTransport(server.handler)
|
|
)
|
|
|
|
|
|
class TestSelfHostedBackend:
|
|
# --- constructor / auth setup (the crux of the bug) -------------------
|
|
|
|
def test_init_uses_x_api_key_not_token_auth(self):
|
|
b = SelfHostedBackend("adminkey", "http://sh:8888")
|
|
assert b._client.headers["x-api-key"] == "adminkey"
|
|
assert "authorization" not in b._client.headers # NOT the cloud 'Token' scheme
|
|
|
|
|
|
# --- search ----------------------------------------------------------
|
|
|
|
|
|
# --- add / update / delete ------------------------------------------
|
|
|
|
|
|
# --- error propagation (feeds the plugin's circuit breaker) ----------
|
|
|
|
def test_http_error_raises(self):
|
|
s = _StubServer()
|
|
with pytest.raises(httpx.HTTPStatusError):
|
|
_backend(s).delete("missing") # 404 -> raise_for_status; 'not found' won't trip breaker
|