pip install fastapi uvicorn openai
cat > /workspace/milvus_rag.py << 'EOF'
from fastapi import FastAPI
from pydantic import BaseModel
from pymilvus import Collection, connections
from sentence_transformers import SentenceTransformer
from openai import OpenAI
import os
app = FastAPI(title="Milvus RAG API")
# 在启动时初始化
connections.connect("default", host="localhost", port="19530")
collection = Collection("documents")
collection.load()
embedder = SentenceTransformer("all-MiniLM-L6-v2", device="cuda")
llm = OpenAI(api_key=os.environ["OPENAI_API_KEY"])
class QueryRequest(BaseModel):
question: str
n_results: int = 5
@app.get("/health")
async def health():
return {"status": "ok", "vectors": collection.num_entities}
@app.post("/search")
async def semantic_search(req: QueryRequest):
embedding = embedder.encode(
[req.question],
normalize_embeddings=True
)[0].tolist()
results = collection.search(
data=[embedding],
anns_field="embedding",
param={"metric_type": "COSINE", "params": {"ef": 64}},
limit=req.n_results,
output_fields=["text", "source", "category"]
)
return {
"results": [
{
"text": hit.entity.get("text"),
"source": hit.entity.get("source"),
"score": hit.score
}
for hit in results[0]
]
}
@app.post("/rag")
async def rag(req: QueryRequest):
embedding = embedder.encode([req.question], normalize_embeddings=True)[0].tolist()
hits = collection.search(
data=[embedding],
anns_field="embedding",
param={"metric_type": "COSINE", "params": {"ef": 64}},
limit=req.n_results,
output_fields=["text", "source"]
)[0]
context = "\n\n".join([
f"[{hit.entity.get('source')}]: {hit.entity.get('text')}"
for hit in hits if hit.score > 0.4
])
response = llm.chat.completions.create(
model="gpt-4o-mini",
messages=[
{"role": "system", "content": "Answer based on context. Be concise."},
{"role": "user", "content": f"Context:\n{context}\n\nQuestion: {req.question}"}
]
)
return {"answer": response.choices[0].message.content, "context_used": len(hits)}
if __name__ == "__main__":
import uvicorn
uvicorn.run(app, host="0.0.0.0", port=8000)
EOF
python3 /workspace/milvus_rag.py