Adding custom connection to weaviate connection client to enable client connection to local container (#3360)

Co-authored-by: Vikram Iyer <vikramiyer@Vikrams-MacBook-Pro.local>
This commit is contained in:
VikramIyer125
2025-08-25 11:24:02 -07:00
committed by GitHub
parent 3b2d0ad0eb
commit 8c8368781d
2 changed files with 62 additions and 60 deletions
+26 -3
View File
@@ -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
+36 -57
View File
@@ -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: