-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathembedding_server.py
More file actions
72 lines (51 loc) · 1.96 KB
/
Copy pathembedding_server.py
File metadata and controls
72 lines (51 loc) · 1.96 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
#!/usr/bin/env python3
"""MemVault Local Embedding Server.
Provides OpenAI-compatible /embeddings endpoint using sentence-transformers.
No API key needed — runs entirely locally.
"""
import logging
import os
from contextlib import asynccontextmanager
from fastapi import FastAPI
from pydantic import BaseModel
from sentence_transformers import SentenceTransformer
logging.basicConfig(level=logging.INFO)
logger = logging.getLogger(__name__)
os.environ.setdefault("no_proxy", "localhost,127.0.0.1")
os.environ.setdefault("NO_PROXY", "localhost,127.0.0.1")
MODEL_NAME = os.environ.get("MEMVAULT_EMBEDDING_MODEL", "all-MiniLM-L6-v2")
PORT = int(os.environ.get("MEMVAULT_EMBEDDING_PORT", "8001"))
model = None
@asynccontextmanager
async def lifespan(app: FastAPI):
global model
logger.info(f"Loading embedding model: {MODEL_NAME}...")
model = SentenceTransformer(MODEL_NAME)
logger.info("Model loaded!")
yield
logger.info("Shutting down...")
app = FastAPI(title="MemVault Embedding Server", lifespan=lifespan)
class EmbedRequest(BaseModel):
input: list[str] | str
model: str = "all-MiniLM-L6-v2"
class EmbedResponse(BaseModel):
data: list[dict]
model: str
usage: dict
@app.post("/embeddings")
@app.post("/v1/embeddings")
async def embed(req: EmbedRequest) -> EmbedResponse:
"""Create embeddings (OpenAI-compatible format)."""
inputs = req.input if isinstance(req.input, list) else [req.input]
embeddings = model.encode(inputs).tolist()
return EmbedResponse(
data=[{"embedding": emb, "index": i, "object": "embedding"} for i, emb in enumerate(embeddings)],
model=MODEL_NAME,
usage={"prompt_tokens": sum(len(t.split()) for t in inputs), "total_tokens": sum(len(t.split()) for t in inputs)},
)
@app.get("/health")
async def health():
return {"status": "ok", "model": MODEL_NAME, "loaded": model is not None}
if __name__ == "__main__":
import uvicorn
uvicorn.run(app, host="0.0.0.0", port=PORT)