From 5b007d66362eff8f54fbcf0a3a8969fcc749f71c Mon Sep 17 00:00:00 2001 From: shamAnimates <145093437+shamAnimates@users.noreply.github.com> Date: Sun, 16 Aug 2026 08:08:01 +0530 Subject: [PATCH 1/4] fix(openai-sdk-python): isolate middleware clients --- .../src/supermemory_openai/middleware.py | 112 ++++- .../tests/test_client_isolation.py | 396 ++++++++++++++++++ 2 files changed, 500 insertions(+), 8 deletions(-) create mode 100644 packages/openai-sdk-python/tests/test_client_isolation.py diff --git a/packages/openai-sdk-python/src/supermemory_openai/middleware.py b/packages/openai-sdk-python/src/supermemory_openai/middleware.py index 7c83b90c..33dac1c9 100644 --- a/packages/openai-sdk-python/src/supermemory_openai/middleware.py +++ b/packages/openai-sdk-python/src/supermemory_openai/middleware.py @@ -48,6 +48,29 @@ class SupermemoryProfileSearch: self.search_results: dict[str, Any] = data.get("searchResults", {}) +class _ResourceFacade: + """Delegate an SDK resource while overriding selected attributes.""" + + def __init__( + self, + resource: Any, + lazy_overrides: Optional[dict[str, Any]] = None, + **overrides: Any, + ) -> None: + self._resource = resource + self._lazy_overrides = lazy_overrides or {} + for name, value in overrides.items(): + setattr(self, name, value) + + def __getattr__(self, name: str) -> Any: + factory = self._lazy_overrides.get(name) + if factory is not None: + value = factory() + setattr(self, name, value) + return value + return getattr(self._resource, name) + + async def supermemory_profile_search( container_tag: str, query_text: str, @@ -269,7 +292,14 @@ class SupermemoryOpenAIWrapper: openai_client: Union[OpenAI, AsyncOpenAI], options: OpenAIMiddlewareOptions, ): - self._client: Union[OpenAI, AsyncOpenAI] = openai_client + self._client: Union[OpenAI, AsyncOpenAI] = getattr( + openai_client, + "__supermemory_openai_base_client__", + openai_client, + ) + # A stable attribute also lets wrappers from another module copy or hot reload + # recover the pristine client instead of nesting middleware. + self.__supermemory_openai_base_client__ = self._client self._container_tag: str = options.container_tag self._options: OpenAIMiddlewareOptions = options self._logger: Logger = create_logger(self._options.verbose) @@ -293,8 +323,8 @@ class SupermemoryOpenAIWrapper: f"Failed to initialize Supermemory client: {e}", e ) - # Wrap the chat completions create method - self._wrap_chat_completions() + # Expose isolated resource facades without mutating the supplied client. + self.chat = self._create_chat_facade() def _get_api_key(self) -> str: """Get Supermemory API key from environment.""" @@ -307,16 +337,79 @@ class SupermemoryOpenAIWrapper: ) return api_key - def _wrap_chat_completions(self) -> None: - """Wrap the chat completions create method with memory injection.""" - original_create = self._client.chat.completions.create + def _create_chat_facade(self) -> _ResourceFacade: + """Create isolated chat/completions facades with memory injection.""" + completions_resource = self._client.chat.completions + completions = _ResourceFacade( + completions_resource, + lazy_overrides={ + "with_raw_response": lambda: self._create_completion_variant_facade( + "with_raw_response" + ), + "with_streaming_response": lambda: self._create_completion_variant_facade( + "with_streaming_response" + ), + }, + create=self._create_completion_method(completions_resource.create), + ) + return _ResourceFacade( + self._client.chat, + lazy_overrides={ + "with_raw_response": lambda: self._create_chat_variant_facade( + "with_raw_response" + ), + "with_streaming_response": lambda: self._create_chat_variant_facade( + "with_streaming_response" + ), + }, + completions=completions, + ) + def _create_completion_variant_facade(self, name: str) -> _ResourceFacade: + """Preserve raw and streaming response behavior on isolated facades.""" + resource = getattr(self._client.chat.completions, name) + return _ResourceFacade( + resource, + create=self._create_completion_method(resource.create), + ) + + def _create_chat_variant_facade(self, name: str) -> _ResourceFacade: + """Wrap completions reached through a chat response variant.""" + chat_resource = getattr(self._client.chat, name) + completions_resource = chat_resource.completions + return _ResourceFacade( + chat_resource, + completions=_ResourceFacade( + completions_resource, + create=self._create_completion_method(completions_resource.create), + ), + ) + + def _create_client_variant_facade(self, name: str) -> _ResourceFacade: + """Wrap completions reached through a client response variant.""" + client_resource = getattr(self._client, name) + chat_resource = client_resource.chat + completions_resource = chat_resource.completions + return _ResourceFacade( + client_resource, + chat=_ResourceFacade( + chat_resource, + completions=_ResourceFacade( + completions_resource, + create=self._create_completion_method(completions_resource.create), + ), + ), + ) + + def _create_completion_method(self, original_create: Any) -> Any: + """Wrap one completion create implementation with memory injection.""" if asyncio.iscoroutinefunction(original_create): async def create_with_memory( **kwargs: Any, ) -> Any: return await self._create_with_memory_async(original_create, **kwargs) + else: def create_with_memory( @@ -324,8 +417,7 @@ class SupermemoryOpenAIWrapper: ) -> Any: return self._create_with_memory_sync(original_create, **kwargs) - # Replace the create method with our wrapper - setattr(self._client.chat.completions, "create", create_with_memory) + return create_with_memory async def _create_with_memory_async( self, @@ -616,6 +708,10 @@ class SupermemoryOpenAIWrapper: def __getattr__(self, name: str) -> Any: """Delegate all other attributes to the wrapped client.""" + if name in {"with_raw_response", "with_streaming_response"}: + value = self._create_client_variant_facade(name) + setattr(self, name, value) + return value return getattr(self._client, name) diff --git a/packages/openai-sdk-python/tests/test_client_isolation.py b/packages/openai-sdk-python/tests/test_client_isolation.py new file mode 100644 index 00000000..700ac4cf --- /dev/null +++ b/packages/openai-sdk-python/tests/test_client_isolation.py @@ -0,0 +1,396 @@ +"""Regression tests for shared OpenAI client middleware isolation.""" + +from __future__ import annotations + +import asyncio +import os +from types import SimpleNamespace +from typing import Any, Generator, Literal, Optional +from unittest.mock import AsyncMock, Mock, patch + +import pytest + +from supermemory_openai import OpenAIMiddlewareOptions, with_supermemory + + +def middleware_options( + container_tag: str, + add_memory: Literal["always", "never"] = "never", +) -> OpenAIMiddlewareOptions: + return OpenAIMiddlewareOptions( + container_tag=container_tag, + custom_id=f"thread-{container_tag}", + add_memory=add_memory, + ) + + +def create_sync_client() -> tuple[Any, Mock]: + create = Mock(return_value={"id": "chat-response"}) + completions = SimpleNamespace(create=create, marker="completions-marker") + chat = SimpleNamespace(completions=completions, marker="chat-marker") + return SimpleNamespace(chat=chat), create + + +def create_async_client() -> tuple[Any, AsyncMock]: + create = AsyncMock(return_value={"id": "chat-response"}) + completions = SimpleNamespace(create=create) + chat = SimpleNamespace(completions=completions) + return SimpleNamespace(chat=chat), create + + +class RawResponse: + def __init__(self, label: str) -> None: + self.label = label + + def parse(self) -> str: + return f"parsed-{self.label}" + + +class SyncStreamContext: + def __init__(self, label: str) -> None: + self.label = label + + def __enter__(self) -> "SyncStreamContext": + return self + + def __exit__(self, *args: Any) -> None: + return None + + +class AsyncStreamContext: + def __init__(self, label: str) -> None: + self.label = label + + async def __aenter__(self) -> "AsyncStreamContext": + return self + + async def __aexit__(self, *args: Any) -> None: + return None + + +def attach_response_variants( + client: Any, + raw_create_factory: Any, + streaming_create_factory: Any, +) -> dict[str, Any]: + calls: dict[str, Any] = {} + + def completion_resource(label: str, factory: Any) -> Any: + create = factory(label) + calls[label] = create + return SimpleNamespace(create=create) + + client.chat.completions.with_raw_response = completion_resource( + "completions-raw", raw_create_factory + ) + client.chat.with_raw_response = SimpleNamespace( + completions=completion_resource("chat-raw", raw_create_factory) + ) + client.with_raw_response = SimpleNamespace( + chat=SimpleNamespace( + completions=completion_resource("client-raw", raw_create_factory) + ) + ) + client.chat.completions.with_streaming_response = completion_resource( + "completions-stream", streaming_create_factory + ) + client.chat.with_streaming_response = SimpleNamespace( + completions=completion_resource("chat-stream", streaming_create_factory) + ) + client.with_streaming_response = SimpleNamespace( + chat=SimpleNamespace( + completions=completion_resource("client-stream", streaming_create_factory) + ) + ) + return calls + + +def response_variant_creates(client: Any, name: str) -> list[Any]: + return [ + getattr(client.chat.completions, name).create, + getattr(client.chat, name).completions.create, + getattr(client, name).chat.completions.create, + ] + + +@pytest.fixture(autouse=True) # type: ignore[untyped-decorator] +def supermemory_api_key() -> Generator[None, None, None]: + with patch.dict(os.environ, {"SUPERMEMORY_API_KEY": "test-key"}): + yield + + +def test_shared_sync_client_does_not_stack_tenant_middleware() -> None: + base_client, original_create = create_sync_client() + lookups: list[str] = [] + + async def fake_prompt( + messages: list[Any], + container_tag: str, + logger: Any, + mode: Any, + api_key: str, + ) -> list[Any]: + lookups.append(container_tag) + return [ + {"role": "system", "content": f"secret-{container_tag}"}, + *messages, + ] + + with patch( + "supermemory_openai.middleware.supermemory.Supermemory", + return_value=Mock(), + ), patch( + "supermemory_openai.middleware.add_system_prompt", + side_effect=fake_prompt, + ): + tenant_a: Any = with_supermemory( + base_client, + middleware_options("tenant-a"), + ) + tenant_b: Any = with_supermemory( + base_client, + middleware_options("tenant-b"), + ) + + assert base_client.chat.completions.create is original_create + assert tenant_a.chat is not base_client.chat + assert tenant_b.chat.marker == "chat-marker" + assert tenant_b.chat.completions.marker == "completions-marker" + + tenant_b.chat.completions.create( + model="gpt-test", + messages=[{"role": "user", "content": "private"}], + ) + + assert lookups == ["tenant-b"] + messages = original_create.call_args.kwargs["messages"] + assert messages[0]["content"] == "secret-tenant-b" + assert all("tenant-a" not in str(message) for message in messages) + + base_client.chat.completions.create( + model="gpt-test", + messages=[{"role": "user", "content": "unwrapped"}], + ) + assert lookups == ["tenant-b"] + + +def test_rewrapping_a_facade_recovers_the_pristine_client() -> None: + base_client, original_create = create_sync_client() + lookups: list[str] = [] + + async def fake_prompt( + messages: list[Any], + container_tag: str, + logger: Any, + mode: Any, + api_key: str, + ) -> list[Any]: + lookups.append(container_tag) + return [ + {"role": "system", "content": f"secret-{container_tag}"}, + *messages, + ] + + with patch( + "supermemory_openai.middleware.supermemory.Supermemory", + return_value=Mock(), + ), patch( + "supermemory_openai.middleware.add_system_prompt", + side_effect=fake_prompt, + ): + tenant_a: Any = with_supermemory( + base_client, + middleware_options("tenant-a"), + ) + tenant_b: Any = with_supermemory( + tenant_a, + middleware_options("tenant-b"), + ) + + tenant_b.chat.completions.create( + model="gpt-test", + messages=[{"role": "user", "content": "private"}], + ) + + assert lookups == ["tenant-b"] + assert original_create.call_count == 1 + assert base_client.chat.completions.create is original_create + + +def test_raw_and_streaming_response_facades_remain_memory_aware() -> None: + base_client, _ = create_sync_client() + calls = attach_response_variants( + base_client, + lambda label: Mock(return_value=RawResponse(label)), + lambda label: Mock(return_value=SyncStreamContext(label)), + ) + lookups: list[str] = [] + + async def fake_prompt( + messages: list[Any], + container_tag: str, + logger: Any, + mode: Any, + api_key: str, + ) -> list[Any]: + lookups.append(container_tag) + return [ + {"role": "system", "content": f"secret-{container_tag}"}, + *messages, + ] + + with patch( + "supermemory_openai.middleware.supermemory.Supermemory", + return_value=Mock(), + ), patch( + "supermemory_openai.middleware.add_system_prompt", + side_effect=fake_prompt, + ): + tenant_b: Any = with_supermemory( + base_client, + middleware_options("tenant-b"), + ) + + for create in response_variant_creates(tenant_b, "with_raw_response"): + response = create( + model="gpt-test", + messages=[{"role": "user", "content": "raw"}], + ) + assert response.parse().startswith("parsed-") + + for create in response_variant_creates(tenant_b, "with_streaming_response"): + with create( + model="gpt-test", + messages=[{"role": "user", "content": "streaming"}], + ) as stream: + assert stream.label.endswith("stream") + + assert lookups == ["tenant-b"] * 6 + for create in calls.values(): + assert "secret-tenant-b" in str(create.call_args.kwargs["messages"]) + + +def test_async_raw_and_streaming_response_prefixes_remain_memory_aware() -> None: + base_client, _ = create_async_client() + calls = attach_response_variants( + base_client, + lambda label: AsyncMock(return_value=RawResponse(label)), + lambda label: Mock(return_value=AsyncStreamContext(label)), + ) + lookups: list[str] = [] + + async def fake_prompt( + messages: list[Any], + container_tag: str, + logger: Any, + mode: Any, + api_key: str, + ) -> list[Any]: + lookups.append(container_tag) + return [ + {"role": "system", "content": f"secret-{container_tag}"}, + *messages, + ] + + async def call_raw(create: Any) -> None: + response = await create( + model="gpt-test", + messages=[{"role": "user", "content": "raw"}], + ) + assert response.parse().startswith("parsed-") + + async def consume_stream(stream_context: Any) -> None: + async with stream_context as stream: + assert stream.label.endswith("stream") + + with patch( + "supermemory_openai.middleware.supermemory.Supermemory", + return_value=Mock(), + ), patch( + "supermemory_openai.middleware.add_system_prompt", + side_effect=fake_prompt, + ): + tenant_b: Any = with_supermemory( + base_client, + middleware_options("tenant-b"), + ) + + for create in response_variant_creates(tenant_b, "with_raw_response"): + asyncio.run(call_raw(create)) + + for create in response_variant_creates(tenant_b, "with_streaming_response"): + stream_context = create( + model="gpt-test", + messages=[{"role": "user", "content": "streaming"}], + ) + asyncio.run(consume_stream(stream_context)) + + assert lookups == ["tenant-b"] * 6 + for create in calls.values(): + assert "secret-tenant-b" in str(create.call_args.kwargs["messages"]) + + +@pytest.mark.asyncio # type: ignore[untyped-decorator] +async def test_shared_async_client_saves_only_for_selected_tenant() -> None: + base_client, original_create = create_async_client() + lookups: list[str] = [] + writes: list[tuple[str, Optional[str], str]] = [] + + async def fake_prompt( + messages: list[Any], + container_tag: str, + logger: Any, + mode: Any, + api_key: str, + ) -> list[Any]: + lookups.append(container_tag) + return [ + {"role": "system", "content": f"secret-{container_tag}"}, + *messages, + ] + + async def fake_add_memory( + client: Any, + container_tag: str, + content: str, + custom_id: Optional[str], + logger: Any, + ) -> None: + writes.append((container_tag, custom_id, content)) + + with patch( + "supermemory_openai.middleware.supermemory.Supermemory", + return_value=Mock(), + ), patch( + "supermemory_openai.middleware.add_system_prompt", + side_effect=fake_prompt, + ), patch( + "supermemory_openai.middleware.add_memory_tool", + side_effect=fake_add_memory, + ): + tenant_a: Any = with_supermemory( + base_client, + middleware_options("tenant-a", add_memory="always"), + ) + tenant_b: Any = with_supermemory( + base_client, + middleware_options("tenant-b", add_memory="always"), + ) + + await tenant_b.chat.completions.create( + model="gpt-test", + messages=[{"role": "user", "content": "private tenant B message"}], + ) + await tenant_a.wait_for_background_tasks() + await tenant_b.wait_for_background_tasks() + + assert base_client.chat.completions.create is original_create + assert lookups == ["tenant-b"] + assert writes == [ + ( + "tenant-b", + "conversation:thread-tenant-b", + "User: private tenant B message", + ) + ] + assert original_create.call_count == 1 From 879a4707cd74209d70f0c9ba7d412543f346a3c2 Mon Sep 17 00:00:00 2001 From: abhinav7x94 Date: Sun, 16 Aug 2026 09:07:13 +0530 Subject: [PATCH 2/4] fix(openai-sdk-python): handle async response variants --- .../src/supermemory_openai/middleware.py | 1371 ++++++++--------- 1 file changed, 602 insertions(+), 769 deletions(-) diff --git a/packages/openai-sdk-python/src/supermemory_openai/middleware.py b/packages/openai-sdk-python/src/supermemory_openai/middleware.py index 33dac1c9..c377c38d 100644 --- a/packages/openai-sdk-python/src/supermemory_openai/middleware.py +++ b/packages/openai-sdk-python/src/supermemory_openai/middleware.py @@ -1,769 +1,602 @@ -"""Supermemory middleware for OpenAI clients.""" - -import asyncio -import inspect -import os -from dataclasses import dataclass -from typing import Any, Literal, Optional, Union, cast - -import supermemory -from openai import AsyncOpenAI, OpenAI -from openai.types.chat import ( - ChatCompletionMessageParam, - ChatCompletionSystemMessageParam, -) - -from .exceptions import ( - SupermemoryAPIError, - SupermemoryConfigurationError, - SupermemoryMemoryOperationError, - SupermemoryNetworkError, -) -from .utils import ( - Logger, - convert_profile_to_markdown, - create_logger, - deduplicate_memories, - get_conversation_content, - get_last_user_message, -) - - -@dataclass -class OpenAIMiddlewareOptions: - """Configuration options for OpenAI middleware.""" - - container_tag: str # Required: identifies the user/container - custom_id: str # Required: groups messages into the same document - verbose: bool = False - mode: Literal["profile", "query", "full"] = "profile" - add_memory: Literal["always", "never"] = "always" - - -class SupermemoryProfileSearch: - """Type for Supermemory profile search response.""" - - def __init__(self, data: dict[str, Any]): - self.profile: dict[str, Any] = data.get("profile", {}) - self.search_results: dict[str, Any] = data.get("searchResults", {}) - - -class _ResourceFacade: - """Delegate an SDK resource while overriding selected attributes.""" - - def __init__( - self, - resource: Any, - lazy_overrides: Optional[dict[str, Any]] = None, - **overrides: Any, - ) -> None: - self._resource = resource - self._lazy_overrides = lazy_overrides or {} - for name, value in overrides.items(): - setattr(self, name, value) - - def __getattr__(self, name: str) -> Any: - factory = self._lazy_overrides.get(name) - if factory is not None: - value = factory() - setattr(self, name, value) - return value - return getattr(self._resource, name) - - -async def supermemory_profile_search( - container_tag: str, - query_text: str, - api_key: str, -) -> SupermemoryProfileSearch: - """Search for memories using the SuperMemory profile API.""" - payload = { - "containerTag": container_tag, - } - if query_text: - payload["q"] = query_text - - try: - import aiohttp - - async with aiohttp.ClientSession() as session: - async with session.post( - "https://api.supermemory.ai/v4/profile", - headers={ - "Content-Type": "application/json", - "Authorization": f"Bearer {api_key}", - }, - json=payload, - ) as response: - if not response.ok: - error_text = await response.text() - raise SupermemoryAPIError( - "Supermemory profile search failed", - status_code=response.status, - response_text=error_text, - ) - - data = await response.json() - return SupermemoryProfileSearch(data) - - except ImportError: - # Fallback to requests if aiohttp not available - import requests - - response = requests.post( - "https://api.supermemory.ai/v4/profile", - headers={ - "Content-Type": "application/json", - "Authorization": f"Bearer {api_key}", - }, - json=payload, - ) - - if not response.ok: - raise SupermemoryAPIError( - "Supermemory profile search failed", - status_code=response.status_code, - response_text=response.text, - ) - - return SupermemoryProfileSearch(response.json()) - - -async def add_system_prompt( - messages: list[ChatCompletionMessageParam], - container_tag: str, - logger: Logger, - mode: Literal["profile", "query", "full"], - api_key: str, -) -> list[ChatCompletionMessageParam]: - """Add memory-enhanced system prompts to chat completion messages.""" - system_prompt_exists = any(msg.get("role") == "system" for msg in messages) - - query_text = get_last_user_message(messages) if mode != "profile" else "" - - memories_response = await supermemory_profile_search( - container_tag, query_text, api_key - ) - - profile = memories_response.profile or {} - search_results_data = memories_response.search_results or {} - memory_count_static = len(profile.get("static", [])) - memory_count_dynamic = len(profile.get("dynamic", [])) - memory_count_search = len(search_results_data.get("results", [])) - - logger.info( - "Memory search completed", - { - "container_tag": container_tag, - "memory_count_static": memory_count_static, - "memory_count_dynamic": memory_count_dynamic, - "query_text": query_text[:100] + ("..." if len(query_text) > 100 else ""), - "mode": mode, - }, - ) - - deduplicated = deduplicate_memories( - static=profile.get("static", []), - dynamic=profile.get("dynamic", []), - search_results=search_results_data.get("results", []), - ) - - logger.debug( - "Memory deduplication completed", - { - "static": { - "original": memory_count_static, - "deduplicated": len(deduplicated.static), - }, - "dynamic": { - "original": memory_count_dynamic, - "deduplicated": len(deduplicated.dynamic), - }, - "search_results": { - "original": memory_count_search, - "deduplicated": len(deduplicated.search_results), - }, - }, - ) - - profile_data = "" - if mode != "query": - profile_data = convert_profile_to_markdown( - { - "profile": { - "static": deduplicated.static, - "dynamic": deduplicated.dynamic, - }, - "searchResults": {"results": []}, - } - ) - - search_results_memories = "" - if mode != "profile" and deduplicated.search_results: - search_results_memories = ( - "Search results for user's recent message: \n" - + "\n".join(f"- {memory}" for memory in deduplicated.search_results) - ) - - memories = f"{profile_data}\n{search_results_memories}".strip() - - if memories: - logger.debug( - "Memory content preview", - { - "content": memories, - "full_length": len(memories), - }, - ) - - if not memories: - return messages - - if system_prompt_exists: - logger.debug("Added memories to existing system prompt") - return [ - {**msg, "content": f"{msg.get('content', '')} \n {memories}"} - if msg.get("role") == "system" - else msg - for msg in messages - ] - - logger.debug("System prompt does not exist, created system prompt with memories") - system_message: ChatCompletionSystemMessageParam = { - "role": "system", - "content": memories, - } - return [system_message] + messages - - -async def add_memory_tool( - client: supermemory.Supermemory, - container_tag: str, - content: str, - custom_id: Optional[str], - logger: Logger, -) -> None: - """Add a new memory to the SuperMemory system.""" - try: - add_params = { - "content": content, - "container_tags": [container_tag], - } - if custom_id is not None: - add_params["custom_id"] = custom_id - - # Handle both sync and async supermemory clients - result = client.memories.add(**add_params) - if inspect.isawaitable(result): - response = await result - else: - response = result - - logger.info( - "Memory saved successfully", - { - "container_tag": container_tag, - "custom_id": custom_id, - "content_length": len(content), - "memory_id": response.id, - }, - ) - except (OSError, ConnectionError) as network_error: - logger.error( - "Network error while saving memory", - {"error": str(network_error)}, - ) - raise SupermemoryNetworkError( - "Failed to save memory due to network error", network_error - ) - except Exception as error: - logger.error( - "Error saving memory", - {"error": str(error)}, - ) - raise SupermemoryMemoryOperationError("Failed to save memory", error) - - -class SupermemoryOpenAIWrapper: - """Wrapper for OpenAI client with Supermemory middleware.""" - - def __init__( - self, - openai_client: Union[OpenAI, AsyncOpenAI], - options: OpenAIMiddlewareOptions, - ): - self._client: Union[OpenAI, AsyncOpenAI] = getattr( - openai_client, - "__supermemory_openai_base_client__", - openai_client, - ) - # A stable attribute also lets wrappers from another module copy or hot reload - # recover the pristine client instead of nesting middleware. - self.__supermemory_openai_base_client__ = self._client - self._container_tag: str = options.container_tag - self._options: OpenAIMiddlewareOptions = options - self._logger: Logger = create_logger(self._options.verbose) - - # Track background tasks to ensure they complete - self._background_tasks: set[asyncio.Task] = set() - - if not hasattr(supermemory, "Supermemory"): - raise SupermemoryConfigurationError( - "supermemory package is required but not found", - ImportError("supermemory package not installed"), - ) - - api_key = self._get_api_key() - try: - self._supermemory_client: supermemory.Supermemory = supermemory.Supermemory( - api_key=api_key - ) - except Exception as e: - raise SupermemoryConfigurationError( - f"Failed to initialize Supermemory client: {e}", e - ) - - # Expose isolated resource facades without mutating the supplied client. - self.chat = self._create_chat_facade() - - def _get_api_key(self) -> str: - """Get Supermemory API key from environment.""" - import os - - api_key = os.getenv("SUPERMEMORY_API_KEY") - if not api_key: - raise SupermemoryConfigurationError( - "SUPERMEMORY_API_KEY environment variable is required but not set" - ) - return api_key - - def _create_chat_facade(self) -> _ResourceFacade: - """Create isolated chat/completions facades with memory injection.""" - completions_resource = self._client.chat.completions - completions = _ResourceFacade( - completions_resource, - lazy_overrides={ - "with_raw_response": lambda: self._create_completion_variant_facade( - "with_raw_response" - ), - "with_streaming_response": lambda: self._create_completion_variant_facade( - "with_streaming_response" - ), - }, - create=self._create_completion_method(completions_resource.create), - ) - return _ResourceFacade( - self._client.chat, - lazy_overrides={ - "with_raw_response": lambda: self._create_chat_variant_facade( - "with_raw_response" - ), - "with_streaming_response": lambda: self._create_chat_variant_facade( - "with_streaming_response" - ), - }, - completions=completions, - ) - - def _create_completion_variant_facade(self, name: str) -> _ResourceFacade: - """Preserve raw and streaming response behavior on isolated facades.""" - resource = getattr(self._client.chat.completions, name) - return _ResourceFacade( - resource, - create=self._create_completion_method(resource.create), - ) - - def _create_chat_variant_facade(self, name: str) -> _ResourceFacade: - """Wrap completions reached through a chat response variant.""" - chat_resource = getattr(self._client.chat, name) - completions_resource = chat_resource.completions - return _ResourceFacade( - chat_resource, - completions=_ResourceFacade( - completions_resource, - create=self._create_completion_method(completions_resource.create), - ), - ) - - def _create_client_variant_facade(self, name: str) -> _ResourceFacade: - """Wrap completions reached through a client response variant.""" - client_resource = getattr(self._client, name) - chat_resource = client_resource.chat - completions_resource = chat_resource.completions - return _ResourceFacade( - client_resource, - chat=_ResourceFacade( - chat_resource, - completions=_ResourceFacade( - completions_resource, - create=self._create_completion_method(completions_resource.create), - ), - ), - ) - - def _create_completion_method(self, original_create: Any) -> Any: - """Wrap one completion create implementation with memory injection.""" - if asyncio.iscoroutinefunction(original_create): - - async def create_with_memory( - **kwargs: Any, - ) -> Any: - return await self._create_with_memory_async(original_create, **kwargs) - - else: - - def create_with_memory( - **kwargs: Any, - ) -> Any: - return self._create_with_memory_sync(original_create, **kwargs) - - return create_with_memory - - async def _create_with_memory_async( - self, - original_create: Any, - **kwargs: Any, - ) -> Any: - """Async version of create with memory injection.""" - messages = kwargs.get("messages", []) - - if self._options.add_memory == "always": - user_message = get_last_user_message(messages) - if user_message and user_message.strip(): - content = ( - get_conversation_content(messages) - if self._options.custom_id - else user_message - ) - custom_id = ( - f"conversation:{self._options.custom_id}" - if self._options.custom_id - else None - ) - - # Create background task for memory storage - task = asyncio.create_task( - add_memory_tool( - self._supermemory_client, - self._container_tag, - content, - custom_id, - self._logger, - ) - ) - - # Track the task and set up cleanup - self._background_tasks.add(task) - task.add_done_callback(self._background_tasks.discard) - - # Log any exceptions but don't fail the main request - def handle_task_exception(task_obj): - try: - if task_obj.exception() is not None: - exception = task_obj.exception() - if isinstance( - exception, - (SupermemoryNetworkError, SupermemoryAPIError), - ): - self._logger.warn( - "Background memory storage failed", - { - "error": str(exception), - "type": type(exception).__name__, - }, - ) - else: - self._logger.error( - "Unexpected error in background memory storage", - { - "error": str(exception), - "type": type(exception).__name__, - }, - ) - except asyncio.CancelledError: - self._logger.debug("Memory storage task was cancelled") - - task.add_done_callback(handle_task_exception) - - if self._options.mode != "profile": - user_message = get_last_user_message(messages) - if not user_message: - self._logger.debug("No user message found, skipping memory search") - return await original_create(**kwargs) - - self._logger.info( - "Starting memory search", - { - "container_tag": self._container_tag, - "conversation_id": self._options.custom_id, - "mode": self._options.mode, - }, - ) - - enhanced_messages = await add_system_prompt( - messages, - self._container_tag, - self._logger, - self._options.mode, - self._get_api_key(), - ) - - kwargs["messages"] = enhanced_messages - return await original_create(**kwargs) - - def _create_with_memory_sync( - self, - original_create: Any, - **kwargs: Any, - ) -> Any: - """Sync version of create with memory injection.""" - # For sync clients, we implement a simplified version without background tasks - messages = kwargs.get("messages", []) - - # Handle memory addition synchronously if needed - if self._options.add_memory == "always": - user_message = get_last_user_message(messages) - if user_message and user_message.strip(): - content = ( - get_conversation_content(messages) - if self._options.custom_id - else user_message - ) - custom_id = ( - f"conversation:{self._options.custom_id}" - if self._options.custom_id - else None - ) - - # Use asyncio.run() for the memory addition - try: - asyncio.run( - add_memory_tool( - self._supermemory_client, - self._container_tag, - content, - custom_id, - self._logger, - ) - ) - except RuntimeError as e: - if "cannot be called from a running event loop" in str(e): - # We're in an async context, log warning and skip memory saving - self._logger.warn( - "Cannot save memory in sync client from async context", - {"error": str(e)}, - ) - else: - raise - except SupermemoryNetworkError as e: - # Network errors are expected, log as warning - self._logger.warn("Network error saving memory", {"error": str(e)}) - except (SupermemoryAPIError, SupermemoryMemoryOperationError) as e: - # API/memory errors are concerning, log as error - self._logger.error("Failed to save memory", {"error": str(e)}) - except Exception as e: - # Unexpected errors should be investigated - self._logger.error( - "Unexpected error saving memory", - {"error": str(e), "type": type(e).__name__}, - ) - - # Handle memory search and injection - if self._options.mode != "profile": - user_message = get_last_user_message(messages) - if not user_message: - self._logger.debug("No user message found, skipping memory search") - return original_create(**kwargs) - - self._logger.info( - "Starting memory search", - { - "container_tag": self._container_tag, - "conversation_id": self._options.custom_id, - "mode": self._options.mode, - }, - ) - - # Use asyncio.run() for memory search and injection - try: - enhanced_messages = asyncio.run( - add_system_prompt( - messages, - self._container_tag, - self._logger, - self._options.mode, - self._get_api_key(), - ) - ) - except RuntimeError as e: - if "cannot be called from a running event loop" in str(e): - # We're in an async context, run in a separate thread - import concurrent.futures - - with concurrent.futures.ThreadPoolExecutor(max_workers=1) as executor: - future = executor.submit( - asyncio.run, - add_system_prompt( - messages, - self._container_tag, - self._logger, - self._options.mode, - self._get_api_key(), - ), - ) - enhanced_messages = future.result() - else: - raise - - kwargs["messages"] = enhanced_messages - return original_create(**kwargs) - - async def wait_for_background_tasks(self, timeout: Optional[float] = 10.0) -> None: - """ - Wait for all background memory storage tasks to complete. - - Args: - timeout: Maximum time to wait in seconds. None for no timeout. - - Raises: - asyncio.TimeoutError: If tasks don't complete within timeout - """ - if not self._background_tasks: - return - - self._logger.debug( - f"Waiting for {len(self._background_tasks)} background tasks to complete" - ) - - try: - if timeout is not None: - await asyncio.wait_for( - asyncio.gather(*self._background_tasks, return_exceptions=True), - timeout=timeout, - ) - else: - await asyncio.gather(*self._background_tasks, return_exceptions=True) - - self._logger.debug("All background tasks completed") - except asyncio.TimeoutError: - self._logger.warn( - f"Background tasks did not complete within {timeout}s timeout" - ) - # Cancel remaining tasks - tasks_to_cancel = [task for task in self._background_tasks if not task.done()] - for task in tasks_to_cancel: - task.cancel() - - if tasks_to_cancel: - await asyncio.gather(*tasks_to_cancel, return_exceptions=True) - raise - - def cancel_background_tasks(self) -> None: - """Cancel all pending background tasks.""" - cancelled_count = 0 - for task in self._background_tasks: - if not task.done(): - task.cancel() - cancelled_count += 1 - - if cancelled_count > 0: - self._logger.debug(f"Cancelled {cancelled_count} pending background tasks") - - async def __aenter__(self): - """Async context manager entry.""" - return self - - async def __aexit__(self, exc_type, exc_val, exc_tb): - """Async context manager exit - wait for background tasks.""" - try: - await self.wait_for_background_tasks(timeout=5.0) - except asyncio.TimeoutError: - self._logger.warn("Some background memory tasks did not complete on exit") - - def __enter__(self): - """Sync context manager entry.""" - return self - - def __exit__(self, exc_type, exc_val, exc_tb): - """Sync context manager exit - attempt to wait for background tasks.""" - if self._background_tasks: - try: - # Try to wait for background tasks in sync context - asyncio.run(self.wait_for_background_tasks(timeout=5.0)) - except RuntimeError as e: - if "cannot be called from a running event loop" in str(e): - # In async context, just cancel the tasks - self._logger.warn( - "Cannot wait for background tasks in sync context from async environment. " - "Use async context manager or call wait_for_background_tasks() manually." - ) - self.cancel_background_tasks() - else: - raise - except asyncio.TimeoutError: - self._logger.warn( - "Some background memory tasks did not complete on exit" - ) - self.cancel_background_tasks() - - def __getattr__(self, name: str) -> Any: - """Delegate all other attributes to the wrapped client.""" - if name in {"with_raw_response", "with_streaming_response"}: - value = self._create_client_variant_facade(name) - setattr(self, name, value) - return value - return getattr(self._client, name) - - -def with_supermemory( - openai_client: Union[OpenAI, AsyncOpenAI], - options: OpenAIMiddlewareOptions, -) -> Union[OpenAI, AsyncOpenAI]: - """ - Wraps an OpenAI client with SuperMemory middleware to automatically inject relevant memories - into the system prompt based on the user's message content. - - This middleware searches the supermemory API for relevant memories using the container tag - and user message, then either appends memories to an existing system prompt or creates - a new system prompt with the memories. - - Args: - openai_client: The OpenAI client to wrap with SuperMemory middleware - options: Configuration options for the middleware (container_tag and custom_id are required) - - Returns: - An OpenAI client with SuperMemory middleware injected - - Example: - ```python - from supermemory_openai import with_supermemory, OpenAIMiddlewareOptions - from openai import OpenAI - - # Create OpenAI client with supermemory middleware - openai = OpenAI(api_key=os.getenv("OPENAI_API_KEY")) - openai_with_supermemory = with_supermemory( - openai, - OpenAIMiddlewareOptions( - container_tag="user-123", - custom_id="conversation-456", - mode="full", - add_memory="always" - ) - ) - - # Use normally - memories will be automatically injected - response = await openai_with_supermemory.chat.completions.create( - model="gpt-4", - messages=[ - {"role": "user", "content": "What's my favorite programming language?"} - ] - ) - ``` - - Raises: - ValueError: When SUPERMEMORY_API_KEY environment variable is not set - Exception: When supermemory API request fails - """ - wrapper = SupermemoryOpenAIWrapper(openai_client, options) - # Return the wrapper, which delegates all attributes to the original client - return cast(Union[OpenAI, AsyncOpenAI], wrapper) +"""Supermemory middleware for OpenAI clients.""" + +import asyncio +import inspect +import os +from dataclasses import dataclass +from typing import Any, Literal, Optional, Union, cast + +import supermemory +from openai import AsyncOpenAI, OpenAI +from openai.types.chat import ( + ChatCompletionMessageParam, + ChatCompletionSystemMessageParam, +) + +from .exceptions import ( + SupermemoryAPIError, + SupermemoryConfigurationError, + SupermemoryMemoryOperationError, + SupermemoryNetworkError, +) +from .utils import ( + Logger, + convert_profile_to_markdown, + create_logger, + deduplicate_memories, + get_conversation_content, + get_last_user_message, +) + + +@dataclass +class OpenAIMiddlewareOptions: + """Configuration options for OpenAI middleware.""" + + container_tag: str # Required: identifies the user/container + custom_id: str # Required: groups messages into the same document + verbose: bool = False + mode: Literal["profile", "query", "full"] = "profile" + add_memory: Literal["always", "never"] = "always" + + +class SupermemoryProfileSearch: + """Type for Supermemory profile search response.""" + + def __init__(self, data: dict[str, Any]): + self.profile: dict[str, Any] = data.get("profile", {}) + self.search_results: dict[str, Any] = data.get("searchResults", {}) + + +class _ResourceFacade: + """Delegate an SDK resource while overriding selected attributes.""" + + def __init__( + self, + resource: Any, + lazy_overrides: Optional[dict[str, Any]] = None, + **overrides: Any, + ) -> None: + self._resource = resource + self._lazy_overrides = lazy_overrides or {} + for name, value in overrides.items(): + setattr(self, name, value) + + def __getattr__(self, name: str) -> Any: + factory = self._lazy_overrides.get(name) + if factory is not None: + value = factory() + setattr(self, name, value) + return value + return getattr(self._resource, name) + + +class _AsyncMemoryResponseContextManager: + """Apply async middleware before entering a streaming response context.""" + + def __init__( + self, + wrapper: "SupermemoryOpenAIWrapper", + original_create: Any, + kwargs: dict[str, Any], + ) -> None: + self._wrapper = wrapper + self._original_create = original_create + self._kwargs = kwargs + self._response_context: Any = None + + async def __aenter__(self) -> Any: + async def enter_original(**kwargs: Any) -> Any: + self._response_context = self._original_create(**kwargs) + return await self._response_context.__aenter__() + + return await self._wrapper._create_with_memory_async( + enter_original, + **self._kwargs, + ) + + async def __aexit__( + self, + exc_type: Any, + exc_value: Any, + traceback: Any, + ) -> Any: + if self._response_context is None: + return None + return await self._response_context.__aexit__( + exc_type, + exc_value, + traceback, + ) + + +async def supermemory_profile_search( + container_tag: str, + query_text: str, + api_key: str, +) -> SupermemoryProfileSearch: + """Search for memories using the SuperMemory profile API.""" + payload = { + "containerTag": container_tag, + } + if query_text: + payload["q"] = query_text + + try: + import aiohttp + + async with aiohttp.ClientSession() as session: + async with session.post( + "https://api.supermemory.ai/v4/profile", + headers={ + "Content-Type": "application/json", + "Authorization": f"Bearer {api_key}", + }, + json=payload, + ) as response: + if not response.ok: + error_text = await response.text() + raise SupermemoryAPIError( + "Supermemory profile search failed", + status_code=response.status, + response_text=error_text, + ) + + data = await response.json() + return SupermemoryProfileSearch(data) + + except ImportError: + # Fallback to requests if aiohttp not available + import requests + + response = requests.post( + "https://api.supermemory.ai/v4/profile", + headers={ + "Content-Type": "application/json", + "Authorization": f"Bearer {api_key}", + }, + json=payload, + ) + + if not response.ok: + raise SupermemoryAPIError( + "Supermemory profile search failed", + status_code=response.status_code, + response_text=response.text, + ) + + return SupermemoryProfileSearch(response.json()) + + +async def add_system_prompt( + messages: list[ChatCompletionMessageParam], + container_tag: str, + logger: Logger, + mode: Literal["profile", "query", "full"], + api_key: str, +) -> list[ChatCompletionMessageParam]: + """Add memory-enhanced system prompts to chat completion messages.""" + system_prompt_exists = any(msg.get("role") == "system" for msg in messages) + + query_text = get_last_user_message(messages) if mode != "profile" else "" + + memories_response = await supermemory_profile_search( + container_tag, query_text, api_key + ) + + profile = memories_response.profile or {} + search_results_data = memories_response.search_results or {} + memory_count_static = len(profile.get("static", [])) + memory_count_dynamic = len(profile.get("dynamic", [])) + memory_count_search = len(search_results_data.get("results", [])) + + logger.info( + "Memory search completed", + { + "container_tag": container_tag, + "memory_count_static": memory_count_static, + "memory_count_dynamic": memory_count_dynamic, + "query_text": query_text[:100] + ("..." if len(query_text) > 100 else ""), + "mode": mode, + }, + ) + + deduplicated = deduplicate_memories( + static=profile.get("static", []), + dynamic=profile.get("dynamic", []), + search_results=search_results_data.get("results", []), + ) + + logger.debug( + "Memory deduplication completed", + { + "static": { + "original": memory_count_static, + "deduplicated": len(deduplicated.static), + }, + "dynamic": { + "original": memory_count_dynamic, + "deduplicated": len(deduplicated.dynamic), + }, + "search_results": { + "original": memory_count_search, + "deduplicated": len(deduplicated.search_results), + }, + }, + ) + + profile_data = "" + if mode != "query": + profile_data = convert_profile_to_markdown( + { + "profile": { + "static": deduplicated.static, + "dynamic": deduplicated.dynamic, + }, + "searchResults": {"results": []}, + } + ) + + search_results_memories = "" + if mode != "profile" and deduplicated.search_results: + search_results_memories = ( + "Search results for user's recent message: \n" + + "\n".join(f"- {memory}" for memory in deduplicated.search_results) + ) + + memories = f"{profile_data}\n{search_results_memories}".strip() + + if memories: + logger.debug( + "Memory content preview", + { + "content": memories, + "full_length": len(memories), + }, + ) + + if not memories: + return messages + + if system_prompt_exists: + logger.debug("Added memories to existing system prompt") + return [ + ( + {**msg, "content": f"{msg.get('content', '')} \n {memories}"} + if msg.get("role") == "system" + else msg + ) + for msg in messages + ] + + logger.debug("System prompt does not exist, created system prompt with memories") + system_message: ChatCompletionSystemMessageParam = { + "role": "system", + "content": memories, + } + return [system_message] + messages + + +async def add_memory_tool( + client: supermemory.Supermemory, + container_tag: str, + content: str, + custom_id: Optional[str], + logger: Logger, +) -> None: + """Add a new memory to the SuperMemory system.""" + try: + add_params = { + "content": content, + "container_tags": [container_tag], + } + if custom_id is not None: + add_params["custom_id"] = custom_id + + # Handle both sync and async supermemory clients + result = client.memories.add(**add_params) + if inspect.isawaitable(result): + response = await result + else: + response = result + + logger.info( + "Memory saved successfully", + { + "container_tag": container_tag, + "custom_id": custom_id, + "content_length": len(content), + "memory_id": response.id, + }, + ) + except (OSError, ConnectionError) as network_error: + logger.error( + "Network error while saving memory", + {"error": str(network_error)}, + ) + raise SupermemoryNetworkError( + "Failed to save memory due to network error", network_error + ) + except Exception as error: + logger.error( + "Error saving memory", + {"error": str(error)}, + ) + raise SupermemoryMemoryOperationError("Failed to save memory", error) + + +class SupermemoryOpenAIWrapper: + """Wrapper for OpenAI client with Supermemory middleware.""" + + def __init__( + self, + openai_client: Union[OpenAI, AsyncOpenAI], + options: OpenAIMiddlewareOptions, + ): + self._client: Union[OpenAI, AsyncOpenAI] = getattr( + openai_client, + "__supermemory_openai_base_client__", + openai_client, + ) + # A stable attribute also lets wrappers from another module copy or hot reload + # recover the pristine client instead of nesting middleware. + self.__supermemory_openai_base_client__ = self._client + self._container_tag: str = options.container_tag + self._options: OpenAIMiddlewareOptions = options + self._logger: Logger = create_logger(self._options.verbose) + base_create = self._client.chat.completions.create + self._is_async_client = isinstance( + self._client, AsyncOpenAI + ) or inspect.iscoroutinefunction(inspect.unwrap(base_create)) + + # Track background tasks to ensure they complete + self._background_tasks: set[asyncio.Task] = set() + + if not hasattr(supermemory, "Supermemory"): + raise SupermemoryConfigurationError( + "supermemory package is required but not found", + ImportError("supermemory package not installed"), + ) + + api_key = self._get_api_key() + try: + self._supermemory_client: supermemory.Supermemory = supermemory.Supermemory( + api_key=api_key + ) + except Exception as e: + raise SupermemoryConfigurationError( + f"Failed to initialize Supermemory client: {e}", e + ) + + # Expose isolated resource facades without mutating the supplied client. + self.chat = self._create_chat_facade() + + def _get_api_key(self) -> str: + """Get Supermemory API key from environment.""" + import os + + api_key = os.getenv("SUPERMEMORY_API_KEY") + if not api_key: + raise SupermemoryConfigurationError( + "SUPERMEMORY_API_KEY environment variable is required but not set" + ) + return api_key + + def _create_chat_facade(self) -> _ResourceFacade: + """Create isolated chat/completions facades with memory injection.""" + completions_resource = self._client.chat.completions + completions = _ResourceFacade( + completions_resource, + lazy_overrides={ + "with_raw_response": lambda: self._create_completion_variant_facade( + "with_raw_response" + ), + "with_streaming_response": lambda: self._create_completion_variant_facade( + "with_streaming_response" + ), + }, + create=self._create_completion_method(completions_resource.create), + ) + return _ResourceFacade( + self._client.chat, + lazy_overrides={ + "with_raw_response": lambda: self._create_chat_variant_facade( + "with_raw_response" + ), + "with_streaming_response": lambda: self._create_chat_variant_facade( + "with_streaming_response" + ), + }, + completions=completions, + ) + + def _create_completion_variant_facade(self, name: str) -> _ResourceFacade: + """Preserve raw and streaming response behavior on isolated facades.""" + resource = getattr(self._client.chat.completions, name) + return _ResourceFacade( + resource, + create=self._create_completion_method( + resource.create, + streaming_response=name == "with_streaming_response", + ), + ) + + def _create_chat_variant_facade(self, name: str) -> _ResourceFacade: + """Wrap completions reached through a chat response variant.""" + chat_resource = getattr(self._client.chat, name) + completions_resource = chat_resource.completions + return _ResourceFacade( + chat_resource, + completions=_ResourceFacade( + completions_resource, + create=self._create_completion_method( + completions_resource.create, + streaming_response=name == "with_streaming_response", + ), + ), + ) + + def _create_client_variant_facade(self, name: str) -> _ResourceFacade: + """Wrap completions reached through a client response variant.""" + client_resource = getattr(self._client, name) + chat_resource = client_resource.chaNkwX\][ۜ\]ȋ]ܙX]W٘XܞB +BY[ ] ]ܘ]ܙ\ۜHH[\S[Y\XJ\][ۜX\][ۗܙ\\J] \]ȋ]ܙX]W٘XܞJB +BY[ ]ܘ]ܙ\ۜHH[\S[Y\XJ]T[\S[Y\XJ\][ۜX\][ۗܙ\\JY[ \]ȋ]ܙX]W٘XܞJB +B +BY[ ] \][ۜ˝]X[Z[ܙ\ۜHH\][ۗܙ\\J\][ۜ\X[HX[Z[ܙX]W٘XܞB +BY[ ] ]X[Z[ܙ\ۜHH[\S[Y\XJ\][ۜX\][ۗܙ\\J] \X[HX[Z[ܙX]W٘XܞJB +BY[ ]X[Z[ܙ\ۜHH[\S[Y\XJ]T[\S[Y\XJ\][ۜX\][ۗܙ\\JY[ \X[HX[Z[ܙX]W٘XܞJB +B +B]\[‚Y\ۜWݘ\X[ܙX]\Y[[KNH O\[WN]\ˆ]]Y[ ] \][ۜJKܙX]K]]Y[ ] JK\][ۜ˘ܙX]K]]Y[ JK] \][ۜ˘ܙX]KBPSTSTUSӗUH +ܛX[\][ۜ˜]ȋ] ]ȋY[ ]ȋ\][ۜ˜X[Z[ȋ] X[Z[ȋY[ X[Z[ȋBYX[\[\][ۗܙX]JY[[K]H O[NY]OHܛX[]\Y[ ] \][ۜ˘ܙX]BY]OH\][ۜ˜]Ȏ]\Y[ ] \][ۜ˝]ܘ]ܙ\ۜKܙX]BY]OH] ]Ȏ]\Y[ ] ]ܘ]ܙ\ۜK\][ۜ˘ܙX]BY]OHY[ ]Ȏ]\Y[ ]ܘ]ܙ\ۜK] \][ۜ˘ܙX]BY]OH\][ۜ˜X[Z[Ȏ]\Y[ ] \][ۜ˝]X[Z[ܙ\ۜKܙX]BY]OH] X[Z[Ȏ]\Y[ ] ]X[Z[ܙ\ۜK\][ۜ˘ܙX]BY]OHY[ X[Z[Ȏ]\Y[ ]X[Z[ܙ\ۜK] \][ۜ˘ܙX]BZ\H\\[ۑ\܊[ۛۈ\][ۈ]]HB]\ ^\J]]\OUYJH\NYۛܙV[\Y YXܘ]ܗBY\\Y[[ܞW\W^J +H O[\]ܖӛۙKۙKۙWN]] X +˙[\ۋȔTTQSSԖWTWVH\ Z^HJNZY[Y\\Y[Y[\ۛX[[ZY]\J +H OۙN\WY[ ܚY[[ܙX]HHܙX]W[Y[ + +B\Έ\HHB\[YZW\ +Y\Y\Έ\[WK۝Z[\YΈ\[K[N[K\W^N +H O\[WN\˘\[ +۝Z[\YB]\ˆȜH\[H۝[Xܙ] ^۝Z[\YHK +Y\Y\B]] +\\Y[[ܞW[ZKZY]\K\\Y[[ܞK\\Y[[ܞH]\ݘ[YOS[ +K +K] +\\Y[[ܞW[ZKZY]\KY\[W\YWYXYZW\  +N[[N[HH]\\Y[[ܞJ\WY[ ZY]\W[ۜ[[ XHK +B[[؎[HH]\\Y[[ܞJ\WY[ ZY]\W[ۜ[[ XK +B\\\WY[ ] \][ۜ˘ܙX]H\ܚY[[ܙX]B\\[[K]\\WY[ ]\\[[؋] X\\OH] [X\\\\[[؋] \][ۜ˛X\\OH\][ۜ[X\\[[؋] \][ۜ˘ܙX]J[[H ]\Y\Y\VȜH\\۝[]]HWK +B\\\OHȝ[[ XBY\Y\HܚY[[ܙX]K[\˚\țY\Y\ȗB\\Y\Y\VȘ۝[HOHXܙ] ][[ X\\[ +[[ XH[Y\YJH܈Y\YH[Y\Y\B\WY[ ] \][ۜ˘ܙX]J[[H ]\Y\Y\VȜH\\۝[[ܘ\YWK +B\\\OHȝ[[ XBY\ܙ]ܘ\[W٘XYWܙXݙ\W\[WY[ + +H OۙN\WY[ ܚY[[ܙX]HHܙX]W[Y[ + +B\Έ\HHB\[YZW\ +Y\Y\Έ\[WK۝Z[\YΈ\[K[N[K\W^N +H O\[WN\˘\[ +۝Z[\YB]\ˆȜH\[H۝[Xܙ] ^۝Z[\YHK +Y\Y\B]] +\\Y[[ܞW[ZKZY]\K\\Y[[ܞK\\Y[[ܞH]\ݘ[YOS[ +K +K] +\\Y[[ܞW[ZKZY]\KY\[W\YWYXYZW\  +N[[N[HH]\\Y[[ܞJ\WY[ ZY]\W[ۜ[[ XHK +B[[؎[HH]\\Y[[ܞJ[[KZY]\W[ۜ[[ XK +B[[؋] \][ۜ˘ܙX]J[[H ]\Y\Y\VȜH\\۝[]]HWK +B\\\OHȝ[[ XB\\ܚY[[ܙX]K[[OH B\\\WY[ ] \][ۜ˘ܙX]H\ܚY[[ܙX]BY\ܘ][X[Z[ܙ\ۜW٘XY\ܙ[XZ[Y[[ܞW]\J +H OۙN\WY[ HܙX]W[Y[ + +B[H]Xܙ\ۜWݘ\X[\WY[ [XHX[[]\ݘ[YOT]ԙ\ۜJX[ +JK[XHX[[]\ݘ[YOT[X[P۝^ +X[ +JK +B\Έ\HHB\[YZW\ +Y\Y\Έ\[WK۝Z[\YΈ\[K[N[K\W^N +H O\[WN\˘\[ +۝Z[\YB]\ˆȜH\[H۝[Xܙ] ^۝Z[\YHK +Y\Y\B]] +\\Y[[ܞW[ZKZY]\K\\Y[[ܞK\\Y[[ܞH]\ݘ[YOS[ +K +K] +\\Y[[ܞW[ZKZY]\KY\[W\YWYXYZW\  +N[[؎[HH]\\Y[[ܞJ\WY[ ZY]\W[ۜ[[ XK +B܈ܙX]H[\ۜWݘ\X[ܙX]\[[؋]ܘ]ܙ\ۜHN\ۜHHܙX]J[[H ]\Y\Y\VȜH\\۝[]ȟWK +B\\\ۜK\J +K\] +\Y HB܈ܙX]H[\ۜWݘ\X[ܙX]\[[؋]X[Z[ܙ\ۜHN]ܙX]J[[H ]\Y\Y\VȜH\\۝[X[Z[ȟWK +H\X[N\\X[KX[ [] +X[HB\\\OHȝ[[ XH + ܈ܙX]H[[˝[Y\ +N\\Xܙ] ][[ X[ܙX]K[\˚\țY\Y\ȗJBY\\[ܘ][X[Z[ܙ\ۜWY^\ܙ[XZ[Y[[ܞW]\J +H OۙN\WY[ HܙX]W\[Y[ + +B[H]Xܙ\ۜWݘ\X[\WY[ [XHX[\[[]\ݘ[YOT]ԙ\ۜJX[ +JK[XHX[[]\ݘ[YOP\[X[P۝^ +X[ +JK +B\Έ\HHB\[YZW\ +Y\Y\Έ\[WK۝Z[\YΈ\[K[N[K\W^N +H O\[WN\˘\[ +۝Z[\YB]\ˆȜH\[H۝[Xܙ] ^۝Z[\YHK +Y\Y\B\[Y[ܘ]ܙX]N[JH OۙN\ۜHH]Z]ܙX]J[[H ]\Y\Y\VȜH\\۝[]ȟWK +B\\\ۜK\J +K\] +\Y HB\[Yۜ[YWX[JX[W۝^[JH OۙN\[]X[W۝^\X[N\\X[KX[ [] +X[HB]] +\\Y[[ܞW[ZKZY]\K\\Y[[ܞK\\Y[[ܞH]\ݘ[YOS[ +K +K] +\\Y[[ܞW[ZKZY]\KY\[W\YWYXYZW\  +N[[؎[HH]\\Y[[ܞJ\WY[ ZY]\W[ۜ[[ XK +B܈ܙX]H[\ۜWݘ\X[ܙX]\[[؋]ܘ]ܙ\ۜHN\[[˜[[ܘ]ܙX]JJB܈ܙX]H[\ۜWݘ\X[ܙX]\[[؋]X[Z[ܙ\ۜHNX[W۝^HܙX]J[[H ]\Y\Y\VȜH\\۝[X[Z[ȟWK +B\[[˜[ۜ[YWX[JX[W۝^ +JB\\\OHȝ[[ XH + ܈ܙX]H[[˝[Y\ +N\\Xܙ] ][[ X[ܙX]K[\˚\țY\Y\ȗJB]\ X\˜\[Y]^J]PSTSTUSӗUB]\ X\˘\[[˜\[Y\ܙX[\[[ZW]\W\[ZY]\J]H OۙN\Έ\HHBܚ]\Έ\\V[ۘ[KWHHB[Y\Y\Έ\\[WWHHB\[YZW\ +Y\Y\Έ\[WK۝Z[\YΈ\[K[N[K\W^N +H O\[WN\˘\[ +۝Z[\YB]\ˆȜH\[H۝[Xܙ] ^۝Z[\YHK +Y\Y\B\[YZWYY[[ܞJY[[K۝Z[\YΈ۝[\WY[ۘ[K\[K +H OۙNܚ]\˘\[ + +۝Z[\Y\WY ۝[ +JBY[Wܙ\]Y\ +\]Y\ \]Y\ +H O \ۜNHHۋY\]Y\ ۝[ +B[Y\Y\˘\[ +VțY\Y\ȗJB]\ \ۜJ \]Y\\\]Y\ XY\^Ș۝[ ]\H\X][ۋڜۈKۏ^ˆY]\ ]\ؚX] \][ۈܙX]Y [[ ]\X\Ȏˆˆ[^ Y\YHȜH\\[۝[ȟK[\ܙX\ۈBKK +BY[H \[Y[ +[ܝZ [[ܝ +[Wܙ\]Y\ +JB\WY[H\[[RJ\W^OH[ZK]\Y[ZY[ +BN]] +\\Y[[ܞW[ZKZY]\K\\Y[[ܞK\\Y[[ܞH]\ݘ[YOS[ +K +K] +\\Y[[ܞW[ZKZY]\KY\[W\]YZW\  +K] +\\Y[[ܞW[ZKZY]\KYY[[ܞW]YZWYY[[ܞK +K\[˘]\[XܙUYB +H\]Y\[˜[\Y[\[^\ȋ[[YU\[Bܘ\Y[HH]\\Y[[ܞJ\WY[ ZY]\W[ۜ[[ \X[YY[[ܞOH[^\ȊK +BܙX]HHX[\[\][ۗܙX]Jܘ\Y ] +B\Hˆ[[ ]\Y\Y\ȎȜH\\۝[]]HY\YHWKBY]OHܛX[\ۜHH]Z]ܙX]J +\B\\\ۜKYOH]\ ]\[Y] [] +]ȊN]ܙ\ۜHH]Z]ܙX]J +\B\\]ܙ\ۜK\J +KYOH]\ ]\[N\ۜW۝^HܙX]J +\B\\[X \]Z]XJ\ۜW۝^ +B\[]\ۜW۝^\X[Z[ܙ\ۜN\\X[Z[ܙ\ۜK]\HOH ]Z]ܘ\Y Z]ٛܗؘXܛ[\ +B]Z]\[[˜Y\ + +B˘X + +B[[YW\[Hˆ\[ˆ܈\[[]YY\X\\[˘]YܞK[[YU\[BB\\\OHȝ[[ \X[B\\ܚ]\OHˆ +[[ \X[۝\][ێXY ][[ \X[\\]]HY\YH +BB\\[[Y\Y\HOH B\\[Y\Y\VVȘ۝[HOHXܙ] ][[ \X[\\[[YW\[OHB[[N]Z]\WY[ J +B]\ X\˘\[[\NYۛܙV[\Y YXܘ]ܗB\[Y\\Y\[Y[]\ۛWٛܗ[XY[[ + +H OۙN\WY[ ܚY[[ܙX]HHܙX]W\[Y[ + +B\Έ\HHBܚ]\Έ\\V[ۘ[KWHHB\[YZW\ +Y\Y\Έ\[WK۝Z[\YΈ\[K[N[K\W^N +H O\[WN\˘\[ +۝Z[\YB]\ˆȜH\[H۝[Xܙ] ^۝Z[\YHK +Y\Y\B\[YZWYY[[ܞJY[[K۝Z[\YΈ۝[\WY[ۘ[K\[K +H OۙNܚ]\˘\[ + +۝Z[\Y\WY ۝[ +JB]] +\\Y[[ܞW[ZKZY]\K\\Y[[ܞK\\Y[[ܞH]\ݘ[YOS[ +K +K] +\\Y[[ܞW[ZKZY]\KY\[W\YWYXYZW\  +K] +\\Y[[ܞW[ZKZY]\KYY[[ܞWYWYXYZWYY[[ܞK +N[[N[HH]\\Y[[ܞJ\WY[ ZY]\W[ۜ[[ XHYY[[ܞOH[^\ȊK +B[[؎[HH]\\Y[[ܞJ\WY[ ZY]\W[ۜ[[ XYY[[ܞOH[^\ȊK +B]Z][[؋] \][ۜ˘ܙX]J[[H ]\Y\Y\VȜH\\۝[]]H[[Y\YHWK +B]Z][[KZ]ٛܗؘXܛ[\ +B]Z][[؋Z]ٛܗؘXܛ[\ +B\\\WY[ ] \][ۜ˘ܙX]H\ܚY[[ܙX]B\\\OHȝ[[ XB\\ܚ]\OHˆ +[[ X۝\][ێXY ][[ X\\]]H[[Y\YH +BB\\ܚY[[ܙX]K[[OH B \ No newline at end of file From 25ebd8ca81c078489b5a1e11f7026f84d272673a Mon Sep 17 00:00:00 2001 From: abhinav7x94 Date: Sun, 16 Aug 2026 09:07:55 +0530 Subject: [PATCH 3/4] fix(openai-sdk-python): restore complete async middleware update --- .../src/supermemory_openai/middleware.py | 549 +++++++++++++----- .../tests/test_client_isolation.py | 150 +++++ 2 files changed, 545 insertions(+), 154 deletions(-) diff --git a/packages/openai-sdk-python/src/supermemory_openai/middleware.py b/packages/openai-sdk-python/src/supermemory_openai/middleware.py index c377c38d..01315369 100644 --- a/packages/openai-sdk-python/src/supermemory_openai/middleware.py +++ b/packages/openai-sdk-python/src/supermemory_openai/middleware.py @@ -439,164 +439,405 @@ class SupermemoryOpenAIWrapper: def _create_client_variant_facade(self, name: str) -> _ResourceFacade: """Wrap completions reached through a client response variant.""" client_resource = getattr(self._client, name) - chat_resource = client_resource.chaNkwX\][ۜ\]ȋ]ܙX]W٘XܞB -BY[ ] ]ܘ]ܙ\ۜHH[\S[Y\XJ\][ۜX\][ۗܙ\\J] \]ȋ]ܙX]W٘XܞJB -BY[ ]ܘ]ܙ\ۜHH[\S[Y\XJ]T[\S[Y\XJ\][ۜX\][ۗܙ\\JY[ \]ȋ]ܙX]W٘XܞJB -B -BY[ ] \][ۜ˝]X[Z[ܙ\ۜHH\][ۗܙ\\J\][ۜ\X[HX[Z[ܙX]W٘XܞB -BY[ ] ]X[Z[ܙ\ۜHH[\S[Y\XJ\][ۜX\][ۗܙ\\J] \X[HX[Z[ܙX]W٘XܞJB -BY[ ]X[Z[ܙ\ۜHH[\S[Y\XJ]T[\S[Y\XJ\][ۜX\][ۗܙ\\JY[ \X[HX[Z[ܙX]W٘XܞJB -B -B]\[‚Y\ۜWݘ\X[ܙX]\Y[[KNH O\[WN]\ˆ]]Y[ ] \][ۜJKܙX]K]]Y[ ] JK\][ۜ˘ܙX]K]]Y[ JK] \][ۜ˘ܙX]KBPSTSTUSӗUH -ܛX[\][ۜ˜]ȋ] ]ȋY[ ]ȋ\][ۜ˜X[Z[ȋ] X[Z[ȋY[ X[Z[ȋBYX[\[\][ۗܙX]JY[[K]H O[NY]OHܛX[]\Y[ ] \][ۜ˘ܙX]BY]OH\][ۜ˜]Ȏ]\Y[ ] \][ۜ˝]ܘ]ܙ\ۜKܙX]BY]OH] ]Ȏ]\Y[ ] ]ܘ]ܙ\ۜK\][ۜ˘ܙX]BY]OHY[ ]Ȏ]\Y[ ]ܘ]ܙ\ۜK] \][ۜ˘ܙX]BY]OH\][ۜ˜X[Z[Ȏ]\Y[ ] \][ۜ˝]X[Z[ܙ\ۜKܙX]BY]OH] X[Z[Ȏ]\Y[ ] ]X[Z[ܙ\ۜK\][ۜ˘ܙX]BY]OHY[ X[Z[Ȏ]\Y[ ]X[Z[ܙ\ۜK] \][ۜ˘ܙX]BZ\H\\[ۑ\܊[ۛۈ\][ۈ]]HB]\ ^\J]]\OUYJH\NYۛܙV[\Y YXܘ]ܗBY\\Y[[ܞW\W^J -H O[\]ܖӛۙKۙKۙWN]] X -˙[\ۋȔTTQSSԖWTWVH\ Z^HJNZY[Y\\Y[Y[\ۛX[[ZY]\J -H OۙN\WY[ ܚY[[ܙX]HHܙX]W[Y[ + chat_resource = client_resource.chat + completions_resource = chat_resource.completions + return _ResourceFacade( + client_resource, + chat=_ResourceFacade( + chat_resource, + completions=_ResourceFacade( + completions_resource, + create=self._create_completion_method( + completions_resource.create, + streaming_response=name == "with_streaming_response", + ), + ), + ), + ) + + def _create_completion_method( + self, + original_create: Any, + *, + streaming_response: bool = False, + ) -> Any: + """Wrap one completion create implementation with memory injection.""" + if self._is_async_client and streaming_response: -B\Έ\HHB\[YZW\ -Y\Y\Έ\[WK۝Z[\YΈ\[K[N[K\W^N -H O\[WN\˘\[ -۝Z[\YB]\ˆȜH\[H۝[Xܙ] ^۝Z[\YHK -Y\Y\B]] -\\Y[[ܞW[ZKZY]\K\\Y[[ܞK\\Y[[ܞH]\ݘ[YOS[ -K -K] -\\Y[[ܞW[ZKZY]\KY\[W\YWYXYZW\  -N[[N[HH]\\Y[[ܞJ\WY[ ZY]\W[ۜ[[ XHK -B[[؎[HH]\\Y[[ܞJ\WY[ ZY]\W[ۜ[[ XK -B\\\WY[ ] \][ۜ˘ܙX]H\ܚY[[ܙX]B\\[[K]\\WY[ ]\\[[؋] X\\OH] [X\\\\[[؋] \][ۜ˛X\\OH\][ۜ[X\\[[؋] \][ۜ˘ܙX]J[[H ]\Y\Y\VȜH\\۝[]]HWK -B\\\OHȝ[[ XBY\Y\HܚY[[ܙX]K[\˚\țY\Y\ȗB\\Y\Y\VȘ۝[HOHXܙ] ][[ X\\[ -[[ XH[Y\YJH܈Y\YH[Y\Y\B\WY[ ] \][ۜ˘ܙX]J[[H ]\Y\Y\VȜH\\۝[[ܘ\YWK -B\\\OHȝ[[ XBY\ܙ]ܘ\[W٘XYWܙXݙ\W\[WY[ + def create_streaming_with_memory( + **kwargs: Any, + ) -> _AsyncMemoryResponseContextManager: + return _AsyncMemoryResponseContextManager( + self, + original_create, + kwargs, + ) -H OۙN\WY[ ܚY[[ܙX]HHܙX]W[Y[ + return create_streaming_with_memory -B\Έ\HHB\[YZW\ -Y\Y\Έ\[WK۝Z[\YΈ\[K[N[K\W^N -H O\[WN\˘\[ -۝Z[\YB]\ˆȜH\[H۝[Xܙ] ^۝Z[\YHK -Y\Y\B]] -\\Y[[ܞW[ZKZY]\K\\Y[[ܞK\\Y[[ܞH]\ݘ[YOS[ -K -K] -\\Y[[ܞW[ZKZY]\KY\[W\YWYXYZW\  -N[[N[HH]\\Y[[ܞJ\WY[ ZY]\W[ۜ[[ XHK -B[[؎[HH]\\Y[[ܞJ[[KZY]\W[ۜ[[ XK -B[[؋] \][ۜ˘ܙX]J[[H ]\Y\Y\VȜH\\۝[]]HWK -B\\\OHȝ[[ XB\\ܚY[[ܙX]K[[OH B\\\WY[ ] \][ۜ˘ܙX]H\ܚY[[ܙX]BY\ܘ][X[Z[ܙ\ۜW٘XY\ܙ[XZ[Y[[ܞW]\J -H OۙN\WY[ HܙX]W[Y[ + if self._is_async_client: -B[H]Xܙ\ۜWݘ\X[\WY[ [XHX[[]\ݘ[YOT]ԙ\ۜJX[ -JK[XHX[[]\ݘ[YOT[X[P۝^ -X[ -JK -B\Έ\HHB\[YZW\ -Y\Y\Έ\[WK۝Z[\YΈ\[K[N[K\W^N -H O\[WN\˘\[ -۝Z[\YB]\ˆȜH\[H۝[Xܙ] ^۝Z[\YHK -Y\Y\B]] -\\Y[[ܞW[ZKZY]\K\\Y[[ܞK\\Y[[ܞH]\ݘ[YOS[ -K -K] -\\Y[[ܞW[ZKZY]\KY\[W\YWYXYZW\  -N[[؎[HH]\\Y[[ܞJ\WY[ ZY]\W[ۜ[[ XK -B܈ܙX]H[\ۜWݘ\X[ܙX]\[[؋]ܘ]ܙ\ۜHN\ۜHHܙX]J[[H ]\Y\Y\VȜH\\۝[]ȟWK -B\\\ۜK\J -K\] -\Y HB܈ܙX]H[\ۜWݘ\X[ܙX]\[[؋]X[Z[ܙ\ۜHN]ܙX]J[[H ]\Y\Y\VȜH\\۝[X[Z[ȟWK -H\X[N\\X[KX[ [] -X[HB\\\OHȝ[[ XH - ܈ܙX]H[[˝[Y\ -N\\Xܙ] ][[ X[ܙX]K[\˚\țY\Y\ȗJBY\\[ܘ][X[Z[ܙ\ۜWY^\ܙ[XZ[Y[[ܞW]\J -H OۙN\WY[ HܙX]W\[Y[ + async def create_async_with_memory( + **kwargs: Any, + ) -> Any: + return await self._create_with_memory_async(original_create, **kwargs) -B[H]Xܙ\ۜWݘ\X[\WY[ [XHX[\[[]\ݘ[YOT]ԙ\ۜJX[ -JK[XHX[[]\ݘ[YOP\[X[P۝^ -X[ -JK -B\Έ\HHB\[YZW\ -Y\Y\Έ\[WK۝Z[\YΈ\[K[N[K\W^N -H O\[WN\˘\[ -۝Z[\YB]\ˆȜH\[H۝[Xܙ] ^۝Z[\YHK -Y\Y\B\[Y[ܘ]ܙX]N[JH OۙN\ۜHH]Z]ܙX]J[[H ]\Y\Y\VȜH\\۝[]ȟWK -B\\\ۜK\J -K\] -\Y HB\[Yۜ[YWX[JX[W۝^[JH OۙN\[]X[W۝^\X[N\\X[KX[ [] -X[HB]] -\\Y[[ܞW[ZKZY]\K\\Y[[ܞK\\Y[[ܞH]\ݘ[YOS[ -K -K] -\\Y[[ܞW[ZKZY]\KY\[W\YWYXYZW\  -N[[؎[HH]\\Y[[ܞJ\WY[ ZY]\W[ۜ[[ XK -B܈ܙX]H[\ۜWݘ\X[ܙX]\[[؋]ܘ]ܙ\ۜHN\[[˜[[ܘ]ܙX]JJB܈ܙX]H[\ۜWݘ\X[ܙX]\[[؋]X[Z[ܙ\ۜHNX[W۝^HܙX]J[[H ]\Y\Y\VȜH\\۝[X[Z[ȟWK -B\[[˜[ۜ[YWX[JX[W۝^ -JB\\\OHȝ[[ XH - ܈ܙX]H[[˝[Y\ -N\\Xܙ] ][[ X[ܙX]K[\˚\țY\Y\ȗJB]\ X\˜\[Y]^J]PSTSTUSӗUB]\ X\˘\[[˜\[Y\ܙX[\[[ZW]\W\[ZY]\J]H OۙN\Έ\HHBܚ]\Έ\\V[ۘ[KWHHB[Y\Y\Έ\\[WWHHB\[YZW\ -Y\Y\Έ\[WK۝Z[\YΈ\[K[N[K\W^N -H O\[WN\˘\[ -۝Z[\YB]\ˆȜH\[H۝[Xܙ] ^۝Z[\YHK -Y\Y\B\[YZWYY[[ܞJY[[K۝Z[\YΈ۝[\WY[ۘ[K\[K -H OۙNܚ]\˘\[ + return create_async_with_memory -۝Z[\Y\WY ۝[ -JBY[Wܙ\]Y\ -\]Y\ \]Y\ -H O \ۜNHHۋY\]Y\ ۝[ -B[Y\Y\˘\[ -VțY\Y\ȗJB]\ \ۜJ \]Y\\\]Y\ XY\^Ș۝[ ]\H\X][ۋڜۈKۏ^ˆY]\ ]\ؚX] \][ۈܙX]Y [[ ]\X\Ȏˆˆ[^ Y\YHȜH\\[۝[ȟK[\ܙX\ۈBKK -BY[H \[Y[ -[ܝZ [[ܝ -[Wܙ\]Y\ -JB\WY[H\[[RJ\W^OH[ZK]\Y[ZY[ -BN]] -\\Y[[ܞW[ZKZY]\K\\Y[[ܞK\\Y[[ܞH]\ݘ[YOS[ -K -K] -\\Y[[ܞW[ZKZY]\KY\[W\]YZW\  -K] -\\Y[[ܞW[ZKZY]\KYY[[ܞW]YZWYY[[ܞK -K\[˘]\[XܙUYB -H\]Y\[˜[\Y[\[^\ȋ[[YU\[Bܘ\Y[HH]\\Y[[ܞJ\WY[ ZY]\W[ۜ[[ \X[YY[[ܞOH[^\ȊK -BܙX]HHX[\[\][ۗܙX]Jܘ\Y ] -B\Hˆ[[ ]\Y\Y\ȎȜH\\۝[]]HY\YHWKBY]OHܛX[\ۜHH]Z]ܙX]J -\B\\\ۜKYOH]\ ]\[Y] [] -]ȊN]ܙ\ۜHH]Z]ܙX]J -\B\\]ܙ\ۜK\J -KYOH]\ ]\[N\ۜW۝^HܙX]J -\B\\[X \]Z]XJ\ۜW۝^ -B\[]\ۜW۝^\X[Z[ܙ\ۜN\\X[Z[ܙ\ۜK]\HOH ]Z]ܘ\Y Z]ٛܗؘXܛ[\ -B]Z]\[[˜Y\ - -B˘X + def create_sync_with_memory( + **kwargs: Any, + ) -> Any: + return self._create_with_memory_sync(original_create, **kwargs) -B[[YW\[Hˆ\[ˆ܈\[[]YY\X\\[˘]YܞK[[YU\[BB\\\OHȝ[[ \X[B\\ܚ]\OHˆ -[[ \X[۝\][ێXY ][[ \X[\\]]HY\YH -BB\\[[Y\Y\HOH B\\[Y\Y\VVȘ۝[HOHXܙ] ][[ \X[\\[[YW\[OHB[[N]Z]\WY[ J -B]\ X\˘\[[\NYۛܙV[\Y YXܘ]ܗB\[Y\\Y\[Y[]\ۛWٛܗ[XY[[ - -H OۙN\WY[ ܚY[[ܙX]HHܙX]W\[Y[ - -B\Έ\HHBܚ]\Έ\\V[ۘ[KWHHB\[YZW\ -Y\Y\Έ\[WK۝Z[\YΈ\[K[N[K\W^N -H O\[WN\˘\[ -۝Z[\YB]\ˆȜH\[H۝[Xܙ] ^۝Z[\YHK -Y\Y\B\[YZWYY[[ܞJY[[K۝Z[\YΈ۝[\WY[ۘ[K\[K -H OۙNܚ]\˘\[ - -۝Z[\Y\WY ۝[ -JB]] -\\Y[[ܞW[ZKZY]\K\\Y[[ܞK\\Y[[ܞH]\ݘ[YOS[ -K -K] -\\Y[[ܞW[ZKZY]\KY\[W\YWYXYZW\  -K] -\\Y[[ܞW[ZKZY]\KYY[[ܞWYWYXYZWYY[[ܞK -N[[N[HH]\\Y[[ܞJ\WY[ ZY]\W[ۜ[[ XHYY[[ܞOH[^\ȊK -B[[؎[HH]\\Y[[ܞJ\WY[ ZY]\W[ۜ[[ XYY[[ܞOH[^\ȊK -B]Z][[؋] \][ۜ˘ܙX]J[[H ]\Y\Y\VȜH\\۝[]]H[[Y\YHWK -B]Z][[KZ]ٛܗؘXܛ[\ -B]Z][[؋Z]ٛܗؘXܛ[\ -B\\\WY[ ] \][ۜ˘ܙX]H\ܚY[[ܙX]B\\\OHȝ[[ XB\\ܚ]\OHˆ -[[ X۝\][ێXY ][[ X\\]]H[[Y\YH -BB\\ܚY[[ܙX]K[[OH B \ No newline at end of file + return create_sync_with_memory + + async def _create_with_memory_async( + self, + original_create: Any, + **kwargs: Any, + ) -> Any: + """Async version of create with memory injection.""" + messages = kwargs.get("messages", []) + + if self._options.add_memory == "always": + user_message = get_last_user_message(messages) + if user_message and user_message.strip(): + content = ( + get_conversation_content(messages) + if self._options.custom_id + else user_message + ) + custom_id = ( + f"conversation:{self._options.custom_id}" + if self._options.custom_id + else None + ) + + # Create background task for memory storage + task = asyncio.create_task( + add_memory_tool( + self._supermemory_client, + self._container_tag, + content, + custom_id, + self._logger, + ) + ) + + # Track the task and set up cleanup + self._background_tasks.add(task) + task.add_done_callback(self._background_tasks.discard) + + # Log any exceptions but don't fail the main request + def handle_task_exception(task_obj): + try: + if task_obj.exception() is not None: + exception = task_obj.exception() + if isinstance( + exception, + (SupermemoryNetworkError, SupermemoryAPIError), + ): + self._logger.warn( + "Background memory storage failed", + { + "error": str(exception), + "type": type(exception).__name__, + }, + ) + else: + self._logger.error( + "Unexpected error in background memory storage", + { + "error": str(exception), + "type": type(exception).__name__, + }, + ) + except asyncio.CancelledError: + self._logger.debug("Memory storage task was cancelled") + + task.add_done_callback(handle_task_exception) + + if self._options.mode != "profile": + user_message = get_last_user_message(messages) + if not user_message: + self._logger.debug("No user message found, skipping memory search") + return await original_create(**kwargs) + + self._logger.info( + "Starting memory search", + { + "container_tag": self._container_tag, + "conversation_id": self._options.custom_id, + "mode": self._options.mode, + }, + ) + + enhanced_messages = await add_system_prompt( + messages, + self._container_tag, + self._logger, + self._options.mode, + self._get_api_key(), + ) + + kwargs["messages"] = enhanced_messages + return await original_create(**kwargs) + + def _create_with_memory_sync( + self, + original_create: Any, + **kwargs: Any, + ) -> Any: + """Sync version of create with memory injection.""" + # For sync clients, we implement a simplified version without background tasks + messages = kwargs.get("messages", []) + + # Handle memory addition synchronously if needed + if self._options.add_memory == "always": + user_message = get_last_user_message(messages) + if user_message and user_message.strip(): + content = ( + get_conversation_content(messages) + if self._options.custom_id + else user_message + ) + custom_id = ( + f"conversation:{self._options.custom_id}" + if self._options.custom_id + else None + ) + + # Use asyncio.run() for the memory addition + try: + asyncio.run( + add_memory_tool( + self._supermemory_client, + self._container_tag, + content, + custom_id, + self._logger, + ) + ) + except RuntimeError as e: + if "cannot be called from a running event loop" in str(e): + # We're in an async context, log warning and skip memory saving + self._logger.warn( + "Cannot save memory in sync client from async context", + {"error": str(e)}, + ) + else: + raise + except SupermemoryNetworkError as e: + # Network errors are expected, log as warning + self._logger.warn("Network error saving memory", {"error": str(e)}) + except (SupermemoryAPIError, SupermemoryMemoryOperationError) as e: + # API/memory errors are concerning, log as error + self._logger.error("Failed to save memory", {"error": str(e)}) + except Exception as e: + # Unexpected errors should be investigated + self._logger.error( + "Unexpected error saving memory", + {"error": str(e), "type": type(e).__name__}, + ) + + # Handle memory search and injection + if self._options.mode != "profile": + user_message = get_last_user_message(messages) + if not user_message: + self._logger.debug("No user message found, skipping memory search") + return original_create(**kwargs) + + self._logger.info( + "Starting memory search", + { + "container_tag": self._container_tag, + "conversation_id": self._options.custom_id, + "mode": self._options.mode, + }, + ) + + # Use asyncio.run() for memory search and injection + try: + enhanced_messages = asyncio.run( + add_system_prompt( + messages, + self._container_tag, + self._logger, + self._options.mode, + self._get_api_key(), + ) + ) + except RuntimeError as e: + if "cannot be called from a running event loop" in str(e): + # We're in an async context, run in a separate thread + import concurrent.futures + + with concurrent.futures.ThreadPoolExecutor(max_workers=1) as executor: + future = executor.submit( + asyncio.run, + add_system_prompt( + messages, + self._container_tag, + self._logger, + self._options.mode, + self._get_api_key(), + ), + ) + enhanced_messages = future.result() + else: + raise + + kwargs["messages"] = enhanced_messages + return original_create(**kwargs) + + async def wait_for_background_tasks(self, timeout: Optional[float] = 10.0) -> None: + """ + Wait for all background memory storage tasks to complete. + + Args: + timeout: Maximum time to wait in seconds. None for no timeout. + + Raises: + asyncio.TimeoutError: If tasks don't complete within timeout + """ + if not self._background_tasks: + return + + self._logger.debug( + f"Waiting for {len(self._background_tasks)} background tasks to complete" + ) + + try: + if timeout is not None: + await asyncio.wait_for( + asyncio.gather(*self._background_tasks, return_exceptions=True), + timeout=timeout, + ) + else: + await asyncio.gather(*self._background_tasks, return_exceptions=True) + + self._logger.debug("All background tasks completed") + except asyncio.TimeoutError: + self._logger.warn( + f"Background tasks did not complete within {timeout}s timeout" + ) + # Cancel remaining tasks + tasks_to_cancel = [ + task for task in self._background_tasks if not task.done() + ] + for task in tasks_to_cancel: + task.cancel() + + if tasks_to_cancel: + await asyncio.gather(*tasks_to_cancel, return_exceptions=True) + raise + + def cancel_background_tasks(self) -> None: + """Cancel all pending background tasks.""" + cancelled_count = 0 + for task in self._background_tasks: + if not task.done(): + task.cancel() + cancelled_count += 1 + + if cancelled_count > 0: + self._logger.debug(f"Cancelled {cancelled_count} pending background tasks") + + async def __aenter__(self): + """Async context manager entry.""" + return self + + async def __aexit__(self, exc_type, exc_val, exc_tb): + """Async context manager exit - wait for background tasks.""" + try: + await self.wait_for_background_tasks(timeout=5.0) + except asyncio.TimeoutError: + self._logger.warn("Some background memory tasks did not complete on exit") + + def __enter__(self): + """Sync context manager entry.""" + return self + + def __exit__(self, exc_type, exc_val, exc_tb): + """Sync context manager exit - attempt to wait for background tasks.""" + if self._background_tasks: + try: + # Try to wait for background tasks in sync context + asyncio.run(self.wait_for_background_tasks(timeout=5.0)) + except RuntimeError as e: + if "cannot be called from a running event loop" in str(e): + # In async context, just cancel the tasks + self._logger.warn( + "Cannot wait for background tasks in sync context from async environment. " + "Use async context manager or call wait_for_background_tasks() manually." + ) + self.cancel_background_tasks() + else: + raise + except asyncio.TimeoutError: + self._logger.warn( + "Some background memory tasks did not complete on exit" + ) + self.cancel_background_tasks() + + def __getattr__(self, name: str) -> Any: + """Delegate all other attributes to the wrapped client.""" + if name in {"with_raw_response", "with_streaming_response"}: + value = self._create_client_variant_facade(name) + setattr(self, name, value) + return value + return getattr(self._client, name) + + +def with_supermemory( + openai_client: Union[OpenAI, AsyncOpenAI], + options: OpenAIMiddlewareOptions, +) -> Union[OpenAI, AsyncOpenAI]: + """ + Wraps an OpenAI client with SuperMemory middleware to automatically inject relevant memories + into the system prompt based on the user's message content. + + This middleware searches the supermemory API for relevant memories using the container tag + and user message, then either appends memories to an existing system prompt or creates + a new system prompt with the memories. + + Args: + openai_client: The OpenAI client to wrap with SuperMemory middleware + options: Configuration options for the middleware (container_tag and custom_id are required) + + Returns: + An OpenAI client with SuperMemory middleware injected + + Example: + ```python + from supermemory_openai import with_supermemory, OpenAIMiddlewareOptions + from openai import OpenAI + + # Create OpenAI client with supermemory middleware + openai = OpenAI(api_key=os.getenv("OPENAI_API_KEY")) + openai_with_supermemory = with_supermemory( + openai, + OpenAIMiddlewareOptions( + container_tag="user-123", + custom_id="conversation-456", + mode="full", + add_memory="always" + ) + ) + + # Use normally - memories will be automatically injected + response = await openai_with_supermemory.chat.completions.create( + model="gpt-4", + messages=[ + {"role": "user", "content": "What's my favorite programming language?"} + ] + ) + ``` + + Raises: + ValueError: When SUPERMEMORY_API_KEY environment variable is not set + Exception: When supermemory API request fails + """ + wrapper = SupermemoryOpenAIWrapper(openai_client, options) + # Return the wrapper, which delegates all attributes to the original client + return cast(Union[OpenAI, AsyncOpenAI], wrapper) diff --git a/packages/openai-sdk-python/tests/test_client_isolation.py b/packages/openai-sdk-python/tests/test_client_isolation.py index 700ac4cf..a0a07ace 100644 --- a/packages/openai-sdk-python/tests/test_client_isolation.py +++ b/packages/openai-sdk-python/tests/test_client_isolation.py @@ -3,12 +3,18 @@ from __future__ import annotations import asyncio +import gc +import inspect +import json import os +import warnings from types import SimpleNamespace from typing import Any, Generator, Literal, Optional from unittest.mock import AsyncMock, Mock, patch +import httpx import pytest +from openai import AsyncOpenAI from supermemory_openai import OpenAIMiddlewareOptions, with_supermemory @@ -113,6 +119,35 @@ def response_variant_creates(client: Any, name: str) -> list[Any]: ] +REAL_ASYNC_COMPLETION_PATHS = ( + "normal", + "completions.raw", + "chat.raw", + "client.raw", + "completions.streaming", + "chat.streaming", + "client.streaming", +) + + +def real_async_completion_create(client: Any, path: str) -> Any: + if path == "normal": + return client.chat.completions.create + if path == "completions.raw": + return client.chat.completions.with_raw_response.create + if path == "chat.raw": + return client.chat.with_raw_response.completions.create + if path == "client.raw": + return client.with_raw_response.chat.completions.create + if path == "completions.streaming": + return client.chat.completions.with_streaming_response.create + if path == "chat.streaming": + return client.chat.with_streaming_response.completions.create + if path == "client.streaming": + return client.with_streaming_response.chat.completions.create + raise AssertionError(f"Unknown completion path: {path}") + + @pytest.fixture(autouse=True) # type: ignore[untyped-decorator] def supermemory_api_key() -> Generator[None, None, None]: with patch.dict(os.environ, {"SUPERMEMORY_API_KEY": "test-key"}): @@ -330,6 +365,121 @@ def test_async_raw_and_streaming_response_prefixes_remain_memory_aware() -> None assert "secret-tenant-b" in str(create.call_args.kwargs["messages"]) +@pytest.mark.parametrize("path", REAL_ASYNC_COMPLETION_PATHS) +@pytest.mark.asyncio +async def test_real_async_openai_paths_use_async_middleware(path: str) -> None: + lookups: list[str] = [] + writes: list[tuple[str, Optional[str], str]] = [] + sent_messages: list[list[Any]] = [] + + async def fake_prompt( + messages: list[Any], + container_tag: str, + logger: Any, + mode: Any, + api_key: str, + ) -> list[Any]: + lookups.append(container_tag) + return [ + {"role": "system", "content": f"secret-{container_tag}"}, + *messages, + ] + + async def fake_add_memory( + client: Any, + container_tag: str, + content: str, + custom_id: Optional[str], + logger: Any, + ) -> None: + writes.append((container_tag, custom_id, content)) + + def handle_request(request: httpx.Request) -> httpx.Response: + body = json.loads(request.content) + sent_messages.append(body["messages"]) + return httpx.Response( + 200, + request=request, + headers={"content-type": "application/json"}, + json={ + "id": "chatcmpl-test", + "object": "chat.completion", + "created": 0, + "model": "gpt-test", + "choices": [ + { + "index": 0, + "message": {"role": "assistant", "content": "ok"}, + "finish_reason": "stop", + } + ], + }, + ) + + http_client = httpx.AsyncClient(transport=httpx.MockTransport(handle_request)) + base_client = AsyncOpenAI(api_key="openai-test", http_client=http_client) + + try: + with patch( + "supermemory_openai.middleware.supermemory.Supermemory", + return_value=Mock(), + ), patch( + "supermemory_openai.middleware.add_system_prompt", + new=fake_prompt, + ), patch( + "supermemory_openai.middleware.add_memory_tool", + new=fake_add_memory, + ), warnings.catch_warnings( + record=True + ) as caught: + warnings.simplefilter("always", RuntimeWarning) + wrapped: Any = with_supermemory( + base_client, + middleware_options("tenant-real", add_memory="always"), + ) + create = real_async_completion_create(wrapped, path) + kwargs = { + "model": "gpt-test", + "messages": [{"role": "user", "content": "private message"}], + } + + if path == "normal": + response = await create(**kwargs) + assert response.id == "chatcmpl-test" + elif path.endswith(".raw"): + raw_response = await create(**kwargs) + assert raw_response.parse().id == "chatcmpl-test" + else: + response_context = create(**kwargs) + assert not inspect.isawaitable(response_context) + async with response_context as streaming_response: + assert streaming_response.status_code == 200 + + await wrapped.wait_for_background_tasks() + await asyncio.sleep(0) + gc.collect() + + runtime_warnings = [ + warning + for warning in caught + if issubclass(warning.category, RuntimeWarning) + ] + + assert lookups == ["tenant-real"] + assert writes == [ + ( + "tenant-real", + "conversation:thread-tenant-real", + "User: private message", + ) + ] + assert len(sent_messages) == 1 + assert sent_messages[0][0]["content"] == "secret-tenant-real" + assert runtime_warnings == [] + finally: + await base_client.close() + + @pytest.mark.asyncio # type: ignore[untyped-decorator] async def test_shared_async_client_saves_only_for_selected_tenant() -> None: base_client, original_create = create_async_client() From 9ab2da9b45aabea615aa7f6d5f5e691f81d3160a Mon Sep 17 00:00:00 2001 From: abhinav7x94 Date: Sun, 16 Aug 2026 09:08:36 +0530 Subject: [PATCH 4/4] fix(openai-sdk-python): restore reviewed async middleware blobs --- .../src/supermemory_openai/middleware.py | 1624 ++++++++--------- 1 file changed, 812 insertions(+), 812 deletions(-) diff --git a/packages/openai-sdk-python/src/supermemory_openai/middleware.py b/packages/openai-sdk-python/src/supermemory_openai/middleware.py index 01315369..f1c4aab8 100644 --- a/packages/openai-sdk-python/src/supermemory_openai/middleware.py +++ b/packages/openai-sdk-python/src/supermemory_openai/middleware.py @@ -1,464 +1,464 @@ -"""Supermemory middleware for OpenAI clients.""" - -import asyncio -import inspect -import os -from dataclasses import dataclass -from typing import Any, Literal, Optional, Union, cast - -import supermemory -from openai import AsyncOpenAI, OpenAI -from openai.types.chat import ( - ChatCompletionMessageParam, - ChatCompletionSystemMessageParam, -) - -from .exceptions import ( - SupermemoryAPIError, - SupermemoryConfigurationError, - SupermemoryMemoryOperationError, - SupermemoryNetworkError, -) -from .utils import ( - Logger, - convert_profile_to_markdown, - create_logger, - deduplicate_memories, - get_conversation_content, - get_last_user_message, -) - - -@dataclass -class OpenAIMiddlewareOptions: - """Configuration options for OpenAI middleware.""" - - container_tag: str # Required: identifies the user/container - custom_id: str # Required: groups messages into the same document - verbose: bool = False - mode: Literal["profile", "query", "full"] = "profile" - add_memory: Literal["always", "never"] = "always" - - -class SupermemoryProfileSearch: - """Type for Supermemory profile search response.""" - - def __init__(self, data: dict[str, Any]): - self.profile: dict[str, Any] = data.get("profile", {}) - self.search_results: dict[str, Any] = data.get("searchResults", {}) - - -class _ResourceFacade: - """Delegate an SDK resource while overriding selected attributes.""" - - def __init__( - self, - resource: Any, - lazy_overrides: Optional[dict[str, Any]] = None, - **overrides: Any, - ) -> None: - self._resource = resource - self._lazy_overrides = lazy_overrides or {} - for name, value in overrides.items(): - setattr(self, name, value) - - def __getattr__(self, name: str) -> Any: - factory = self._lazy_overrides.get(name) - if factory is not None: - value = factory() - setattr(self, name, value) - return value - return getattr(self._resource, name) - - -class _AsyncMemoryResponseContextManager: - """Apply async middleware before entering a streaming response context.""" - - def __init__( - self, - wrapper: "SupermemoryOpenAIWrapper", - original_create: Any, - kwargs: dict[str, Any], - ) -> None: - self._wrapper = wrapper - self._original_create = original_create - self._kwargs = kwargs - self._response_context: Any = None - - async def __aenter__(self) -> Any: - async def enter_original(**kwargs: Any) -> Any: - self._response_context = self._original_create(**kwargs) - return await self._response_context.__aenter__() - - return await self._wrapper._create_with_memory_async( - enter_original, - **self._kwargs, - ) - - async def __aexit__( - self, - exc_type: Any, - exc_value: Any, - traceback: Any, - ) -> Any: - if self._response_context is None: - return None - return await self._response_context.__aexit__( - exc_type, - exc_value, - traceback, - ) - - -async def supermemory_profile_search( - container_tag: str, - query_text: str, - api_key: str, -) -> SupermemoryProfileSearch: - """Search for memories using the SuperMemory profile API.""" - payload = { - "containerTag": container_tag, - } - if query_text: - payload["q"] = query_text - - try: - import aiohttp - - async with aiohttp.ClientSession() as session: - async with session.post( - "https://api.supermemory.ai/v4/profile", - headers={ - "Content-Type": "application/json", - "Authorization": f"Bearer {api_key}", - }, - json=payload, - ) as response: - if not response.ok: - error_text = await response.text() - raise SupermemoryAPIError( - "Supermemory profile search failed", - status_code=response.status, - response_text=error_text, - ) - - data = await response.json() - return SupermemoryProfileSearch(data) - - except ImportError: - # Fallback to requests if aiohttp not available - import requests - - response = requests.post( - "https://api.supermemory.ai/v4/profile", - headers={ - "Content-Type": "application/json", - "Authorization": f"Bearer {api_key}", - }, - json=payload, - ) - - if not response.ok: - raise SupermemoryAPIError( - "Supermemory profile search failed", - status_code=response.status_code, - response_text=response.text, - ) - - return SupermemoryProfileSearch(response.json()) - - -async def add_system_prompt( - messages: list[ChatCompletionMessageParam], - container_tag: str, - logger: Logger, - mode: Literal["profile", "query", "full"], - api_key: str, -) -> list[ChatCompletionMessageParam]: - """Add memory-enhanced system prompts to chat completion messages.""" - system_prompt_exists = any(msg.get("role") == "system" for msg in messages) - - query_text = get_last_user_message(messages) if mode != "profile" else "" - - memories_response = await supermemory_profile_search( - container_tag, query_text, api_key - ) - - profile = memories_response.profile or {} - search_results_data = memories_response.search_results or {} - memory_count_static = len(profile.get("static", [])) - memory_count_dynamic = len(profile.get("dynamic", [])) - memory_count_search = len(search_results_data.get("results", [])) - - logger.info( - "Memory search completed", - { - "container_tag": container_tag, - "memory_count_static": memory_count_static, - "memory_count_dynamic": memory_count_dynamic, - "query_text": query_text[:100] + ("..." if len(query_text) > 100 else ""), - "mode": mode, - }, - ) - - deduplicated = deduplicate_memories( - static=profile.get("static", []), - dynamic=profile.get("dynamic", []), - search_results=search_results_data.get("results", []), - ) - - logger.debug( - "Memory deduplication completed", - { - "static": { - "original": memory_count_static, - "deduplicated": len(deduplicated.static), - }, - "dynamic": { - "original": memory_count_dynamic, - "deduplicated": len(deduplicated.dynamic), - }, - "search_results": { - "original": memory_count_search, - "deduplicated": len(deduplicated.search_results), - }, - }, - ) - - profile_data = "" - if mode != "query": - profile_data = convert_profile_to_markdown( - { - "profile": { - "static": deduplicated.static, - "dynamic": deduplicated.dynamic, - }, - "searchResults": {"results": []}, - } - ) - - search_results_memories = "" - if mode != "profile" and deduplicated.search_results: - search_results_memories = ( - "Search results for user's recent message: \n" - + "\n".join(f"- {memory}" for memory in deduplicated.search_results) - ) - - memories = f"{profile_data}\n{search_results_memories}".strip() - - if memories: - logger.debug( - "Memory content preview", - { - "content": memories, - "full_length": len(memories), - }, - ) - - if not memories: - return messages - - if system_prompt_exists: - logger.debug("Added memories to existing system prompt") - return [ - ( - {**msg, "content": f"{msg.get('content', '')} \n {memories}"} - if msg.get("role") == "system" - else msg - ) - for msg in messages - ] - - logger.debug("System prompt does not exist, created system prompt with memories") - system_message: ChatCompletionSystemMessageParam = { - "role": "system", - "content": memories, - } - return [system_message] + messages - - -async def add_memory_tool( - client: supermemory.Supermemory, - container_tag: str, - content: str, - custom_id: Optional[str], - logger: Logger, -) -> None: - """Add a new memory to the SuperMemory system.""" - try: - add_params = { - "content": content, - "container_tags": [container_tag], - } - if custom_id is not None: - add_params["custom_id"] = custom_id - - # Handle both sync and async supermemory clients - result = client.memories.add(**add_params) - if inspect.isawaitable(result): - response = await result - else: - response = result - - logger.info( - "Memory saved successfully", - { - "container_tag": container_tag, - "custom_id": custom_id, - "content_length": len(content), - "memory_id": response.id, - }, - ) - except (OSError, ConnectionError) as network_error: - logger.error( - "Network error while saving memory", - {"error": str(network_error)}, - ) - raise SupermemoryNetworkError( - "Failed to save memory due to network error", network_error - ) - except Exception as error: - logger.error( - "Error saving memory", - {"error": str(error)}, - ) - raise SupermemoryMemoryOperationError("Failed to save memory", error) - - -class SupermemoryOpenAIWrapper: - """Wrapper for OpenAI client with Supermemory middleware.""" - - def __init__( - self, - openai_client: Union[OpenAI, AsyncOpenAI], - options: OpenAIMiddlewareOptions, - ): - self._client: Union[OpenAI, AsyncOpenAI] = getattr( - openai_client, - "__supermemory_openai_base_client__", - openai_client, - ) - # A stable attribute also lets wrappers from another module copy or hot reload - # recover the pristine client instead of nesting middleware. - self.__supermemory_openai_base_client__ = self._client - self._container_tag: str = options.container_tag - self._options: OpenAIMiddlewareOptions = options - self._logger: Logger = create_logger(self._options.verbose) - base_create = self._client.chat.completions.create - self._is_async_client = isinstance( - self._client, AsyncOpenAI - ) or inspect.iscoroutinefunction(inspect.unwrap(base_create)) - - # Track background tasks to ensure they complete - self._background_tasks: set[asyncio.Task] = set() - - if not hasattr(supermemory, "Supermemory"): - raise SupermemoryConfigurationError( - "supermemory package is required but not found", - ImportError("supermemory package not installed"), - ) - - api_key = self._get_api_key() - try: - self._supermemory_client: supermemory.Supermemory = supermemory.Supermemory( - api_key=api_key - ) - except Exception as e: - raise SupermemoryConfigurationError( - f"Failed to initialize Supermemory client: {e}", e - ) - - # Expose isolated resource facades without mutating the supplied client. - self.chat = self._create_chat_facade() - - def _get_api_key(self) -> str: - """Get Supermemory API key from environment.""" - import os - - api_key = os.getenv("SUPERMEMORY_API_KEY") - if not api_key: - raise SupermemoryConfigurationError( - "SUPERMEMORY_API_KEY environment variable is required but not set" - ) - return api_key - - def _create_chat_facade(self) -> _ResourceFacade: - """Create isolated chat/completions facades with memory injection.""" - completions_resource = self._client.chat.completions - completions = _ResourceFacade( - completions_resource, - lazy_overrides={ - "with_raw_response": lambda: self._create_completion_variant_facade( - "with_raw_response" - ), - "with_streaming_response": lambda: self._create_completion_variant_facade( - "with_streaming_response" - ), - }, - create=self._create_completion_method(completions_resource.create), - ) - return _ResourceFacade( - self._client.chat, - lazy_overrides={ - "with_raw_response": lambda: self._create_chat_variant_facade( - "with_raw_response" - ), - "with_streaming_response": lambda: self._create_chat_variant_facade( - "with_streaming_response" - ), - }, - completions=completions, - ) - - def _create_completion_variant_facade(self, name: str) -> _ResourceFacade: - """Preserve raw and streaming response behavior on isolated facades.""" - resource = getattr(self._client.chat.completions, name) - return _ResourceFacade( - resource, - create=self._create_completion_method( - resource.create, - streaming_response=name == "with_streaming_response", - ), - ) - - def _create_chat_variant_facade(self, name: str) -> _ResourceFacade: - """Wrap completions reached through a chat response variant.""" - chat_resource = getattr(self._client.chat, name) - completions_resource = chat_resource.completions - return _ResourceFacade( - chat_resource, - completions=_ResourceFacade( - completions_resource, - create=self._create_completion_method( - completions_resource.create, - streaming_response=name == "with_streaming_response", - ), - ), - ) - - def _create_client_variant_facade(self, name: str) -> _ResourceFacade: - """Wrap completions reached through a client response variant.""" - client_resource = getattr(self._client, name) - chat_resource = client_resource.chat - completions_resource = chat_resource.completions - return _ResourceFacade( - client_resource, - chat=_ResourceFacade( - chat_resource, - completions=_ResourceFacade( - completions_resource, - create=self._create_completion_method( - completions_resource.create, - streaming_response=name == "with_streaming_response", - ), - ), - ), - ) - - def _create_completion_method( - self, - original_create: Any, - *, +"""Supermemory middleware for OpenAI clients.""" + +import asyncio +import inspect +import os +from dataclasses import dataclass +from typing import Any, Literal, Optional, Union, cast + +import supermemory +from openai import AsyncOpenAI, OpenAI +from openai.types.chat import ( + ChatCompletionMessageParam, + ChatCompletionSystemMessageParam, +) + +from .exceptions import ( + SupermemoryAPIError, + SupermemoryConfigurationError, + SupermemoryMemoryOperationError, + SupermemoryNetworkError, +) +from .utils import ( + Logger, + convert_profile_to_markdown, + create_logger, + deduplicate_memories, + get_conversation_content, + get_last_user_message, +) + + +@dataclass +class OpenAIMiddlewareOptions: + """Configuration options for OpenAI middleware.""" + + container_tag: str # Required: identifies the user/container + custom_id: str # Required: groups messages into the same document + verbose: bool = False + mode: Literal["profile", "query", "full"] = "profile" + add_memory: Literal["always", "never"] = "always" + + +class SupermemoryProfileSearch: + """Type for Supermemory profile search response.""" + + def __init__(self, data: dict[str, Any]): + self.profile: dict[str, Any] = data.get("profile", {}) + self.search_results: dict[str, Any] = data.get("searchResults", {}) + + +class _ResourceFacade: + """Delegate an SDK resource while overriding selected attributes.""" + + def __init__( + self, + resource: Any, + lazy_overrides: Optional[dict[str, Any]] = None, + **overrides: Any, + ) -> None: + self._resource = resource + self._lazy_overrides = lazy_overrides or {} + for name, value in overrides.items(): + setattr(self, name, value) + + def __getattr__(self, name: str) -> Any: + factory = self._lazy_overrides.get(name) + if factory is not None: + value = factory() + setattr(self, name, value) + return value + return getattr(self._resource, name) + + +class _AsyncMemoryResponseContextManager: + """Apply async middleware before entering a streaming response context.""" + + def __init__( + self, + wrapper: "SupermemoryOpenAIWrapper", + original_create: Any, + kwargs: dict[str, Any], + ) -> None: + self._wrapper = wrapper + self._original_create = original_create + self._kwargs = kwargs + self._response_context: Any = None + + async def __aenter__(self) -> Any: + async def enter_original(**kwargs: Any) -> Any: + self._response_context = self._original_create(**kwargs) + return await self._response_context.__aenter__() + + return await self._wrapper._create_with_memory_async( + enter_original, + **self._kwargs, + ) + + async def __aexit__( + self, + exc_type: Any, + exc_value: Any, + traceback: Any, + ) -> Any: + if self._response_context is None: + return None + return await self._response_context.__aexit__( + exc_type, + exc_value, + traceback, + ) + + +async def supermemory_profile_search( + container_tag: str, + query_text: str, + api_key: str, +) -> SupermemoryProfileSearch: + """Search for memories using the SuperMemory profile API.""" + payload = { + "containerTag": container_tag, + } + if query_text: + payload["q"] = query_text + + try: + import aiohttp + + async with aiohttp.ClientSession() as session: + async with session.post( + "https://api.supermemory.ai/v4/profile", + headers={ + "Content-Type": "application/json", + "Authorization": f"Bearer {api_key}", + }, + json=payload, + ) as response: + if not response.ok: + error_text = await response.text() + raise SupermemoryAPIError( + "Supermemory profile search failed", + status_code=response.status, + response_text=error_text, + ) + + data = await response.json() + return SupermemoryProfileSearch(data) + + except ImportError: + # Fallback to requests if aiohttp not available + import requests + + response = requests.post( + "https://api.supermemory.ai/v4/profile", + headers={ + "Content-Type": "application/json", + "Authorization": f"Bearer {api_key}", + }, + json=payload, + ) + + if not response.ok: + raise SupermemoryAPIError( + "Supermemory profile search failed", + status_code=response.status_code, + response_text=response.text, + ) + + return SupermemoryProfileSearch(response.json()) + + +async def add_system_prompt( + messages: list[ChatCompletionMessageParam], + container_tag: str, + logger: Logger, + mode: Literal["profile", "query", "full"], + api_key: str, +) -> list[ChatCompletionMessageParam]: + """Add memory-enhanced system prompts to chat completion messages.""" + system_prompt_exists = any(msg.get("role") == "system" for msg in messages) + + query_text = get_last_user_message(messages) if mode != "profile" else "" + + memories_response = await supermemory_profile_search( + container_tag, query_text, api_key + ) + + profile = memories_response.profile or {} + search_results_data = memories_response.search_results or {} + memory_count_static = len(profile.get("static", [])) + memory_count_dynamic = len(profile.get("dynamic", [])) + memory_count_search = len(search_results_data.get("results", [])) + + logger.info( + "Memory search completed", + { + "container_tag": container_tag, + "memory_count_static": memory_count_static, + "memory_count_dynamic": memory_count_dynamic, + "query_text": query_text[:100] + ("..." if len(query_text) > 100 else ""), + "mode": mode, + }, + ) + + deduplicated = deduplicate_memories( + static=profile.get("static", []), + dynamic=profile.get("dynamic", []), + search_results=search_results_data.get("results", []), + ) + + logger.debug( + "Memory deduplication completed", + { + "static": { + "original": memory_count_static, + "deduplicated": len(deduplicated.static), + }, + "dynamic": { + "original": memory_count_dynamic, + "deduplicated": len(deduplicated.dynamic), + }, + "search_results": { + "original": memory_count_search, + "deduplicated": len(deduplicated.search_results), + }, + }, + ) + + profile_data = "" + if mode != "query": + profile_data = convert_profile_to_markdown( + { + "profile": { + "static": deduplicated.static, + "dynamic": deduplicated.dynamic, + }, + "searchResults": {"results": []}, + } + ) + + search_results_memories = "" + if mode != "profile" and deduplicated.search_results: + search_results_memories = ( + "Search results for user's recent message: \n" + + "\n".join(f"- {memory}" for memory in deduplicated.search_results) + ) + + memories = f"{profile_data}\n{search_results_memories}".strip() + + if memories: + logger.debug( + "Memory content preview", + { + "content": memories, + "full_length": len(memories), + }, + ) + + if not memories: + return messages + + if system_prompt_exists: + logger.debug("Added memories to existing system prompt") + return [ + ( + {**msg, "content": f"{msg.get('content', '')} \n {memories}"} + if msg.get("role") == "system" + else msg + ) + for msg in messages + ] + + logger.debug("System prompt does not exist, created system prompt with memories") + system_message: ChatCompletionSystemMessageParam = { + "role": "system", + "content": memories, + } + return [system_message] + messages + + +async def add_memory_tool( + client: supermemory.Supermemory, + container_tag: str, + content: str, + custom_id: Optional[str], + logger: Logger, +) -> None: + """Add a new memory to the SuperMemory system.""" + try: + add_params = { + "content": content, + "container_tags": [container_tag], + } + if custom_id is not None: + add_params["custom_id"] = custom_id + + # Handle both sync and async supermemory clients + result = client.memories.add(**add_params) + if inspect.isawaitable(result): + response = await result + else: + response = result + + logger.info( + "Memory saved successfully", + { + "container_tag": container_tag, + "custom_id": custom_id, + "content_length": len(content), + "memory_id": response.id, + }, + ) + except (OSError, ConnectionError) as network_error: + logger.error( + "Network error while saving memory", + {"error": str(network_error)}, + ) + raise SupermemoryNetworkError( + "Failed to save memory due to network error", network_error + ) + except Exception as error: + logger.error( + "Error saving memory", + {"error": str(error)}, + ) + raise SupermemoryMemoryOperationError("Failed to save memory", error) + + +class SupermemoryOpenAIWrapper: + """Wrapper for OpenAI client with Supermemory middleware.""" + + def __init__( + self, + openai_client: Union[OpenAI, AsyncOpenAI], + options: OpenAIMiddlewareOptions, + ): + self._client: Union[OpenAI, AsyncOpenAI] = getattr( + openai_client, + "__supermemory_openai_base_client__", + openai_client, + ) + # A stable attribute also lets wrappers from another module copy or hot reload + # recover the pristine client instead of nesting middleware. + self.__supermemory_openai_base_client__ = self._client + self._container_tag: str = options.container_tag + self._options: OpenAIMiddlewareOptions = options + self._logger: Logger = create_logger(self._options.verbose) + base_create = self._client.chat.completions.create + self._is_async_client = isinstance( + self._client, AsyncOpenAI + ) or inspect.iscoroutinefunction(inspect.unwrap(base_create)) + + # Track background tasks to ensure they complete + self._background_tasks: set[asyncio.Task] = set() + + if not hasattr(supermemory, "Supermemory"): + raise SupermemoryConfigurationError( + "supermemory package is required but not found", + ImportError("supermemory package not installed"), + ) + + api_key = self._get_api_key() + try: + self._supermemory_client: supermemory.Supermemory = supermemory.Supermemory( + api_key=api_key + ) + except Exception as e: + raise SupermemoryConfigurationError( + f"Failed to initialize Supermemory client: {e}", e + ) + + # Expose isolated resource facades without mutating the supplied client. + self.chat = self._create_chat_facade() + + def _get_api_key(self) -> str: + """Get Supermemory API key from environment.""" + import os + + api_key = os.getenv("SUPERMEMORY_API_KEY") + if not api_key: + raise SupermemoryConfigurationError( + "SUPERMEMORY_API_KEY environment variable is required but not set" + ) + return api_key + + def _create_chat_facade(self) -> _ResourceFacade: + """Create isolated chat/completions facades with memory injection.""" + completions_resource = self._client.chat.completions + completions = _ResourceFacade( + completions_resource, + lazy_overrides={ + "with_raw_response": lambda: self._create_completion_variant_facade( + "with_raw_response" + ), + "with_streaming_response": lambda: self._create_completion_variant_facade( + "with_streaming_response" + ), + }, + create=self._create_completion_method(completions_resource.create), + ) + return _ResourceFacade( + self._client.chat, + lazy_overrides={ + "with_raw_response": lambda: self._create_chat_variant_facade( + "with_raw_response" + ), + "with_streaming_response": lambda: self._create_chat_variant_facade( + "with_streaming_response" + ), + }, + completions=completions, + ) + + def _create_completion_variant_facade(self, name: str) -> _ResourceFacade: + """Preserve raw and streaming response behavior on isolated facades.""" + resource = getattr(self._client.chat.completions, name) + return _ResourceFacade( + resource, + create=self._create_completion_method( + resource.create, + streaming_response=name == "with_streaming_response", + ), + ) + + def _create_chat_variant_facade(self, name: str) -> _ResourceFacade: + """Wrap completions reached through a chat response variant.""" + chat_resource = getattr(self._client.chat, name) + completions_resource = chat_resource.completions + return _ResourceFacade( + chat_resource, + completions=_ResourceFacade( + completions_resource, + create=self._create_completion_method( + completions_resource.create, + streaming_response=name == "with_streaming_response", + ), + ), + ) + + def _create_client_variant_facade(self, name: str) -> _ResourceFacade: + """Wrap completions reached through a client response variant.""" + client_resource = getattr(self._client, name) + chat_resource = client_resource.chat + completions_resource = chat_resource.completions + return _ResourceFacade( + client_resource, + chat=_ResourceFacade( + chat_resource, + completions=_ResourceFacade( + completions_resource, + create=self._create_completion_method( + completions_resource.create, + streaming_response=name == "with_streaming_response", + ), + ), + ), + ) + + def _create_completion_method( + self, + original_create: Any, + *, streaming_response: bool = False, ) -> Any: """Wrap one completion create implementation with memory injection.""" @@ -490,354 +490,354 @@ class SupermemoryOpenAIWrapper: return self._create_with_memory_sync(original_create, **kwargs) return create_sync_with_memory - - async def _create_with_memory_async( - self, - original_create: Any, - **kwargs: Any, - ) -> Any: - """Async version of create with memory injection.""" - messages = kwargs.get("messages", []) - - if self._options.add_memory == "always": - user_message = get_last_user_message(messages) - if user_message and user_message.strip(): - content = ( - get_conversation_content(messages) - if self._options.custom_id - else user_message - ) - custom_id = ( - f"conversation:{self._options.custom_id}" - if self._options.custom_id - else None - ) - - # Create background task for memory storage - task = asyncio.create_task( - add_memory_tool( - self._supermemory_client, - self._container_tag, - content, - custom_id, - self._logger, - ) - ) - - # Track the task and set up cleanup - self._background_tasks.add(task) - task.add_done_callback(self._background_tasks.discard) - - # Log any exceptions but don't fail the main request - def handle_task_exception(task_obj): - try: - if task_obj.exception() is not None: - exception = task_obj.exception() - if isinstance( - exception, - (SupermemoryNetworkError, SupermemoryAPIError), - ): - self._logger.warn( - "Background memory storage failed", - { - "error": str(exception), - "type": type(exception).__name__, - }, - ) - else: - self._logger.error( - "Unexpected error in background memory storage", - { - "error": str(exception), - "type": type(exception).__name__, - }, - ) - except asyncio.CancelledError: - self._logger.debug("Memory storage task was cancelled") - - task.add_done_callback(handle_task_exception) - - if self._options.mode != "profile": - user_message = get_last_user_message(messages) - if not user_message: - self._logger.debug("No user message found, skipping memory search") - return await original_create(**kwargs) - - self._logger.info( - "Starting memory search", - { - "container_tag": self._container_tag, - "conversation_id": self._options.custom_id, - "mode": self._options.mode, - }, - ) - - enhanced_messages = await add_system_prompt( - messages, - self._container_tag, - self._logger, - self._options.mode, - self._get_api_key(), - ) - - kwargs["messages"] = enhanced_messages - return await original_create(**kwargs) - - def _create_with_memory_sync( - self, - original_create: Any, - **kwargs: Any, - ) -> Any: - """Sync version of create with memory injection.""" - # For sync clients, we implement a simplified version without background tasks - messages = kwargs.get("messages", []) - - # Handle memory addition synchronously if needed - if self._options.add_memory == "always": - user_message = get_last_user_message(messages) - if user_message and user_message.strip(): - content = ( - get_conversation_content(messages) - if self._options.custom_id - else user_message - ) - custom_id = ( - f"conversation:{self._options.custom_id}" - if self._options.custom_id - else None - ) - - # Use asyncio.run() for the memory addition - try: - asyncio.run( - add_memory_tool( - self._supermemory_client, - self._container_tag, - content, - custom_id, - self._logger, - ) - ) - except RuntimeError as e: - if "cannot be called from a running event loop" in str(e): - # We're in an async context, log warning and skip memory saving - self._logger.warn( - "Cannot save memory in sync client from async context", - {"error": str(e)}, - ) - else: - raise - except SupermemoryNetworkError as e: - # Network errors are expected, log as warning - self._logger.warn("Network error saving memory", {"error": str(e)}) - except (SupermemoryAPIError, SupermemoryMemoryOperationError) as e: - # API/memory errors are concerning, log as error - self._logger.error("Failed to save memory", {"error": str(e)}) - except Exception as e: - # Unexpected errors should be investigated - self._logger.error( - "Unexpected error saving memory", - {"error": str(e), "type": type(e).__name__}, - ) - - # Handle memory search and injection - if self._options.mode != "profile": - user_message = get_last_user_message(messages) - if not user_message: - self._logger.debug("No user message found, skipping memory search") - return original_create(**kwargs) - - self._logger.info( - "Starting memory search", - { - "container_tag": self._container_tag, - "conversation_id": self._options.custom_id, - "mode": self._options.mode, - }, - ) - - # Use asyncio.run() for memory search and injection - try: - enhanced_messages = asyncio.run( - add_system_prompt( - messages, - self._container_tag, - self._logger, - self._options.mode, - self._get_api_key(), - ) - ) - except RuntimeError as e: - if "cannot be called from a running event loop" in str(e): - # We're in an async context, run in a separate thread - import concurrent.futures - - with concurrent.futures.ThreadPoolExecutor(max_workers=1) as executor: - future = executor.submit( - asyncio.run, - add_system_prompt( - messages, - self._container_tag, - self._logger, - self._options.mode, - self._get_api_key(), - ), - ) - enhanced_messages = future.result() - else: - raise - - kwargs["messages"] = enhanced_messages - return original_create(**kwargs) - - async def wait_for_background_tasks(self, timeout: Optional[float] = 10.0) -> None: - """ - Wait for all background memory storage tasks to complete. - - Args: - timeout: Maximum time to wait in seconds. None for no timeout. - - Raises: - asyncio.TimeoutError: If tasks don't complete within timeout - """ - if not self._background_tasks: - return - - self._logger.debug( - f"Waiting for {len(self._background_tasks)} background tasks to complete" - ) - - try: - if timeout is not None: - await asyncio.wait_for( - asyncio.gather(*self._background_tasks, return_exceptions=True), - timeout=timeout, - ) - else: - await asyncio.gather(*self._background_tasks, return_exceptions=True) - - self._logger.debug("All background tasks completed") - except asyncio.TimeoutError: - self._logger.warn( - f"Background tasks did not complete within {timeout}s timeout" - ) - # Cancel remaining tasks - tasks_to_cancel = [ - task for task in self._background_tasks if not task.done() - ] - for task in tasks_to_cancel: - task.cancel() - - if tasks_to_cancel: - await asyncio.gather(*tasks_to_cancel, return_exceptions=True) - raise - - def cancel_background_tasks(self) -> None: - """Cancel all pending background tasks.""" - cancelled_count = 0 - for task in self._background_tasks: - if not task.done(): - task.cancel() - cancelled_count += 1 - - if cancelled_count > 0: - self._logger.debug(f"Cancelled {cancelled_count} pending background tasks") - - async def __aenter__(self): - """Async context manager entry.""" - return self - - async def __aexit__(self, exc_type, exc_val, exc_tb): - """Async context manager exit - wait for background tasks.""" - try: - await self.wait_for_background_tasks(timeout=5.0) - except asyncio.TimeoutError: - self._logger.warn("Some background memory tasks did not complete on exit") - - def __enter__(self): - """Sync context manager entry.""" - return self - - def __exit__(self, exc_type, exc_val, exc_tb): - """Sync context manager exit - attempt to wait for background tasks.""" - if self._background_tasks: - try: - # Try to wait for background tasks in sync context - asyncio.run(self.wait_for_background_tasks(timeout=5.0)) - except RuntimeError as e: - if "cannot be called from a running event loop" in str(e): - # In async context, just cancel the tasks - self._logger.warn( - "Cannot wait for background tasks in sync context from async environment. " - "Use async context manager or call wait_for_background_tasks() manually." - ) - self.cancel_background_tasks() - else: - raise - except asyncio.TimeoutError: - self._logger.warn( - "Some background memory tasks did not complete on exit" - ) - self.cancel_background_tasks() - - def __getattr__(self, name: str) -> Any: - """Delegate all other attributes to the wrapped client.""" - if name in {"with_raw_response", "with_streaming_response"}: - value = self._create_client_variant_facade(name) - setattr(self, name, value) - return value - return getattr(self._client, name) - - -def with_supermemory( - openai_client: Union[OpenAI, AsyncOpenAI], - options: OpenAIMiddlewareOptions, -) -> Union[OpenAI, AsyncOpenAI]: - """ - Wraps an OpenAI client with SuperMemory middleware to automatically inject relevant memories - into the system prompt based on the user's message content. - - This middleware searches the supermemory API for relevant memories using the container tag - and user message, then either appends memories to an existing system prompt or creates - a new system prompt with the memories. - - Args: - openai_client: The OpenAI client to wrap with SuperMemory middleware - options: Configuration options for the middleware (container_tag and custom_id are required) - - Returns: - An OpenAI client with SuperMemory middleware injected - - Example: - ```python - from supermemory_openai import with_supermemory, OpenAIMiddlewareOptions - from openai import OpenAI - - # Create OpenAI client with supermemory middleware - openai = OpenAI(api_key=os.getenv("OPENAI_API_KEY")) - openai_with_supermemory = with_supermemory( - openai, - OpenAIMiddlewareOptions( - container_tag="user-123", - custom_id="conversation-456", - mode="full", - add_memory="always" - ) - ) - - # Use normally - memories will be automatically injected - response = await openai_with_supermemory.chat.completions.create( - model="gpt-4", - messages=[ - {"role": "user", "content": "What's my favorite programming language?"} - ] - ) - ``` - - Raises: - ValueError: When SUPERMEMORY_API_KEY environment variable is not set - Exception: When supermemory API request fails - """ - wrapper = SupermemoryOpenAIWrapper(openai_client, options) - # Return the wrapper, which delegates all attributes to the original client - return cast(Union[OpenAI, AsyncOpenAI], wrapper) + + async def _create_with_memory_async( + self, + original_create: Any, + **kwargs: Any, + ) -> Any: + """Async version of create with memory injection.""" + messages = kwargs.get("messages", []) + + if self._options.add_memory == "always": + user_message = get_last_user_message(messages) + if user_message and user_message.strip(): + content = ( + get_conversation_content(messages) + if self._options.custom_id + else user_message + ) + custom_id = ( + f"conversation:{self._options.custom_id}" + if self._options.custom_id + else None + ) + + # Create background task for memory storage + task = asyncio.create_task( + add_memory_tool( + self._supermemory_client, + self._container_tag, + content, + custom_id, + self._logger, + ) + ) + + # Track the task and set up cleanup + self._background_tasks.add(task) + task.add_done_callback(self._background_tasks.discard) + + # Log any exceptions but don't fail the main request + def handle_task_exception(task_obj): + try: + if task_obj.exception() is not None: + exception = task_obj.exception() + if isinstance( + exception, + (SupermemoryNetworkError, SupermemoryAPIError), + ): + self._logger.warn( + "Background memory storage failed", + { + "error": str(exception), + "type": type(exception).__name__, + }, + ) + else: + self._logger.error( + "Unexpected error in background memory storage", + { + "error": str(exception), + "type": type(exception).__name__, + }, + ) + except asyncio.CancelledError: + self._logger.debug("Memory storage task was cancelled") + + task.add_done_callback(handle_task_exception) + + if self._options.mode != "profile": + user_message = get_last_user_message(messages) + if not user_message: + self._logger.debug("No user message found, skipping memory search") + return await original_create(**kwargs) + + self._logger.info( + "Starting memory search", + { + "container_tag": self._container_tag, + "conversation_id": self._options.custom_id, + "mode": self._options.mode, + }, + ) + + enhanced_messages = await add_system_prompt( + messages, + self._container_tag, + self._logger, + self._options.mode, + self._get_api_key(), + ) + + kwargs["messages"] = enhanced_messages + return await original_create(**kwargs) + + def _create_with_memory_sync( + self, + original_create: Any, + **kwargs: Any, + ) -> Any: + """Sync version of create with memory injection.""" + # For sync clients, we implement a simplified version without background tasks + messages = kwargs.get("messages", []) + + # Handle memory addition synchronously if needed + if self._options.add_memory == "always": + user_message = get_last_user_message(messages) + if user_message and user_message.strip(): + content = ( + get_conversation_content(messages) + if self._options.custom_id + else user_message + ) + custom_id = ( + f"conversation:{self._options.custom_id}" + if self._options.custom_id + else None + ) + + # Use asyncio.run() for the memory addition + try: + asyncio.run( + add_memory_tool( + self._supermemory_client, + self._container_tag, + content, + custom_id, + self._logger, + ) + ) + except RuntimeError as e: + if "cannot be called from a running event loop" in str(e): + # We're in an async context, log warning and skip memory saving + self._logger.warn( + "Cannot save memory in sync client from async context", + {"error": str(e)}, + ) + else: + raise + except SupermemoryNetworkError as e: + # Network errors are expected, log as warning + self._logger.warn("Network error saving memory", {"error": str(e)}) + except (SupermemoryAPIError, SupermemoryMemoryOperationError) as e: + # API/memory errors are concerning, log as error + self._logger.error("Failed to save memory", {"error": str(e)}) + except Exception as e: + # Unexpected errors should be investigated + self._logger.error( + "Unexpected error saving memory", + {"error": str(e), "type": type(e).__name__}, + ) + + # Handle memory search and injection + if self._options.mode != "profile": + user_message = get_last_user_message(messages) + if not user_message: + self._logger.debug("No user message found, skipping memory search") + return original_create(**kwargs) + + self._logger.info( + "Starting memory search", + { + "container_tag": self._container_tag, + "conversation_id": self._options.custom_id, + "mode": self._options.mode, + }, + ) + + # Use asyncio.run() for memory search and injection + try: + enhanced_messages = asyncio.run( + add_system_prompt( + messages, + self._container_tag, + self._logger, + self._options.mode, + self._get_api_key(), + ) + ) + except RuntimeError as e: + if "cannot be called from a running event loop" in str(e): + # We're in an async context, run in a separate thread + import concurrent.futures + + with concurrent.futures.ThreadPoolExecutor(max_workers=1) as executor: + future = executor.submit( + asyncio.run, + add_system_prompt( + messages, + self._container_tag, + self._logger, + self._options.mode, + self._get_api_key(), + ), + ) + enhanced_messages = future.result() + else: + raise + + kwargs["messages"] = enhanced_messages + return original_create(**kwargs) + + async def wait_for_background_tasks(self, timeout: Optional[float] = 10.0) -> None: + """ + Wait for all background memory storage tasks to complete. + + Args: + timeout: Maximum time to wait in seconds. None for no timeout. + + Raises: + asyncio.TimeoutError: If tasks don't complete within timeout + """ + if not self._background_tasks: + return + + self._logger.debug( + f"Waiting for {len(self._background_tasks)} background tasks to complete" + ) + + try: + if timeout is not None: + await asyncio.wait_for( + asyncio.gather(*self._background_tasks, return_exceptions=True), + timeout=timeout, + ) + else: + await asyncio.gather(*self._background_tasks, return_exceptions=True) + + self._logger.debug("All background tasks completed") + except asyncio.TimeoutError: + self._logger.warn( + f"Background tasks did not complete within {timeout}s timeout" + ) + # Cancel remaining tasks + tasks_to_cancel = [ + task for task in self._background_tasks if not task.done() + ] + for task in tasks_to_cancel: + task.cancel() + + if tasks_to_cancel: + await asyncio.gather(*tasks_to_cancel, return_exceptions=True) + raise + + def cancel_background_tasks(self) -> None: + """Cancel all pending background tasks.""" + cancelled_count = 0 + for task in self._background_tasks: + if not task.done(): + task.cancel() + cancelled_count += 1 + + if cancelled_count > 0: + self._logger.debug(f"Cancelled {cancelled_count} pending background tasks") + + async def __aenter__(self): + """Async context manager entry.""" + return self + + async def __aexit__(self, exc_type, exc_val, exc_tb): + """Async context manager exit - wait for background tasks.""" + try: + await self.wait_for_background_tasks(timeout=5.0) + except asyncio.TimeoutError: + self._logger.warn("Some background memory tasks did not complete on exit") + + def __enter__(self): + """Sync context manager entry.""" + return self + + def __exit__(self, exc_type, exc_val, exc_tb): + """Sync context manager exit - attempt to wait for background tasks.""" + if self._background_tasks: + try: + # Try to wait for background tasks in sync context + asyncio.run(self.wait_for_background_tasks(timeout=5.0)) + except RuntimeError as e: + if "cannot be called from a running event loop" in str(e): + # In async context, just cancel the tasks + self._logger.warn( + "Cannot wait for background tasks in sync context from async environment. " + "Use async context manager or call wait_for_background_tasks() manually." + ) + self.cancel_background_tasks() + else: + raise + except asyncio.TimeoutError: + self._logger.warn( + "Some background memory tasks did not complete on exit" + ) + self.cancel_background_tasks() + + def __getattr__(self, name: str) -> Any: + """Delegate all other attributes to the wrapped client.""" + if name in {"with_raw_response", "with_streaming_response"}: + value = self._create_client_variant_facade(name) + setattr(self, name, value) + return value + return getattr(self._client, name) + + +def with_supermemory( + openai_client: Union[OpenAI, AsyncOpenAI], + options: OpenAIMiddlewareOptions, +) -> Union[OpenAI, AsyncOpenAI]: + """ + Wraps an OpenAI client with SuperMemory middleware to automatically inject relevant memories + into the system prompt based on the user's message content. + + This middleware searches the supermemory API for relevant memories using the container tag + and user message, then either appends memories to an existing system prompt or creates + a new system prompt with the memories. + + Args: + openai_client: The OpenAI client to wrap with SuperMemory middleware + options: Configuration options for the middleware (container_tag and custom_id are required) + + Returns: + An OpenAI client with SuperMemory middleware injected + + Example: + ```python + from supermemory_openai import with_supermemory, OpenAIMiddlewareOptions + from openai import OpenAI + + # Create OpenAI client with supermemory middleware + openai = OpenAI(api_key=os.getenv("OPENAI_API_KEY")) + openai_with_supermemory = with_supermemory( + openai, + OpenAIMiddlewareOptions( + container_tag="user-123", + custom_id="conversation-456", + mode="full", + add_memory="always" + ) + ) + + # Use normally - memories will be automatically injected + response = await openai_with_supermemory.chat.completions.create( + model="gpt-4", + messages=[ + {"role": "user", "content": "What's my favorite programming language?"} + ] + ) + ``` + + Raises: + ValueError: When SUPERMEMORY_API_KEY environment variable is not set + Exception: When supermemory API request fails + """ + wrapper = SupermemoryOpenAIWrapper(openai_client, options) + # Return the wrapper, which delegates all attributes to the original client + return cast(Union[OpenAI, AsyncOpenAI], wrapper)