diff --git a/mem0/vector_stores/weaviate.py b/mem0/vector_stores/weaviate.py index d5b2fa417..989cc49b1 100644 --- a/mem0/vector_stores/weaviate.py +++ b/mem0/vector_stores/weaviate.py @@ -3,6 +3,7 @@ import uuid from typing import Dict, List, Mapping, Optional from pydantic import BaseModel +from urllib.parse import urlparse try: import weaviate @@ -12,7 +13,7 @@ except ImportError: ) import weaviate.classes.config as wvcc -from weaviate.classes.init import Auth +from weaviate.classes.init import Auth, AdditionalConfig, Timeout from weaviate.classes.query import Filter, MetadataQuery from weaviate.util import get_valid_uuid @@ -47,14 +48,36 @@ class Weaviate(VectorStoreBase): auth_config (dict, optional): Authentication configuration for Weaviate. Defaults to None. additional_headers (dict, optional): Additional headers for requests. Defaults to None. """ - if "localhost" in cluster_url: + if "localhost" in cluster_url: self.client = weaviate.connect_to_local(headers=additional_headers) - else: + elif auth_client_secret: self.client = weaviate.connect_to_wcs( cluster_url=cluster_url, auth_credentials=Auth.api_key(auth_client_secret), headers=additional_headers, ) + else: + parsed = urlparse(cluster_url) # e.g., http://mem0_store:8080 + http_host = parsed.hostname or "localhost" + http_port = parsed.port or (443 if parsed.scheme == "https" else 8080) + http_secure = parsed.scheme == "https" + + # Weaviate gRPC defaults (inside Docker network) + grpc_host = http_host + grpc_port = 50051 + grpc_secure = False + + self.client = weaviate.connect_to_custom( + http_host, + http_port, + http_secure, + grpc_host, + grpc_port, + grpc_secure, + headers=additional_headers, + skip_init_checks=True, + additional_config=AdditionalConfig(timeout=Timeout(init=2.0)) + ) self.collection_name = collection_name self.embedding_model_dims = embedding_model_dims diff --git a/openmemory/api/app/mcp_server.py b/openmemory/api/app/mcp_server.py index 914eb15b3..0911a0dac 100644 --- a/openmemory/api/app/mcp_server.py +++ b/openmemory/api/app/mcp_server.py @@ -31,7 +31,6 @@ from fastapi import FastAPI, Request from fastapi.routing import APIRouter from mcp.server.fastmcp import FastMCP from mcp.server.sse import SseServerTransport -from qdrant_client import models as qdrant_models # Load environment variables load_dotenv() @@ -165,74 +164,54 @@ async def search_memory(query: str) -> str: # Get accessible memory IDs based on ACL user_memories = db.query(Memory).filter(Memory.user_id == user.id).all() accessible_memory_ids = [memory.id for memory in user_memories if check_memory_access_permissions(db, memory, app.id)] - - conditions = [qdrant_models.FieldCondition(key="user_id", match=qdrant_models.MatchValue(value=uid))] - - if accessible_memory_ids: - # Convert UUIDs to strings for Qdrant - accessible_memory_ids_str = [str(memory_id) for memory_id in accessible_memory_ids] - conditions.append(qdrant_models.HasIdCondition(has_id=accessible_memory_ids_str)) - filters = qdrant_models.Filter(must=conditions) + filters = { + "user_id": uid + } + embeddings = memory_client.embedding_model.embed(query, "search") - - hits = memory_client.vector_store.client.query_points( - collection_name=memory_client.vector_store.collection_name, - query=embeddings, - query_filter=filters, - limit=10, + + hits = memory_client.vector_store.search( + query=query, + vectors=embeddings, + limit=10, + filters=filters, ) - # Process search results - memories = hits.points - memories = [ - { - "id": memory.id, - "memory": memory.payload["data"], - "hash": memory.payload.get("hash"), - "created_at": memory.payload.get("created_at"), - "updated_at": memory.payload.get("updated_at"), - "score": memory.score, - } - for memory in memories - ] + allowed = set(str(mid) for mid in accessible_memory_ids) if accessible_memory_ids else None - # Log memory access for each memory found - if isinstance(memories, dict) and 'results' in memories: - print(f"Memories: {memories}") - for memory_data in memories['results']: - if 'id' in memory_data: - memory_id = uuid.UUID(memory_data['id']) - # Create access log entry - access_log = MemoryAccessLog( - memory_id=memory_id, - app_id=app.id, - access_type="search", - metadata_={ - "query": query, - "score": memory_data.get('score'), - "hash": memory_data.get('hash') - } - ) - db.add(access_log) - db.commit() - else: - for memory in memories: - memory_id = uuid.UUID(memory['id']) - # Create access log entry + results = [] + for h in hits: + # All vector db search functions return OutputData class + id, score, payload = h.id, h.score, h.payload + if allowed and h.id is None or h.id not in allowed: + continue + + results.append({ + "id": id, + "memory": payload.get("data"), + "hash": payload.get("hash"), + "created_at": payload.get("created_at"), + "updated_at": payload.get("updated_at"), + "score": score, + }) + + for r in results: + if r.get("id"): access_log = MemoryAccessLog( - memory_id=memory_id, + memory_id=uuid.UUID(r["id"]), app_id=app.id, access_type="search", metadata_={ "query": query, - "score": memory.get('score'), - "hash": memory.get('hash') - } + "score": r.get("score"), + "hash": r.get("hash"), + }, ) db.add(access_log) - db.commit() - return json.dumps(memories, indent=2) + db.commit() + + return json.dumps({"results": results}, indent=2) finally: db.close() except Exception as e: