Import AITURK IDE 1.0.0-beta.1 from Hermes 63279301; preserve MIT license
This commit is contained in:
@@ -0,0 +1,680 @@
|
||||
"""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
|
||||
Reference in New Issue
Block a user