embeding-ai / app.py
chingu27's picture
Update app.py
9b13477 verified
Raw
History Blame Contribute Delete
8.71 kB
# ─────────────────────────────────────────────────────────────────────────────
# ChinguBot AI Server β€” HuggingFace Spaces
#
# Endpoints:
# POST /embed β€” embed query (user message)
# POST /embed-passage β€” embed passage (dokumen/intent)
# POST /embed-batch β€” batch embed
# GET /health β€” health check
# GET /ping β€” keep-alive (cegah HF Spaces tidur)
#
# UPGRADE LOG:
# πŸš€ [v1] Init β€” BGE-M3 via FlagEmbedding
# πŸš€ [v2] Ganti FlagEmbedding β†’ sentence-transformers (lebih stabil)
# ─────────────────────────────────────────────────────────────────────────────
import os
import time
import logging
from contextlib import asynccontextmanager
from typing import List, Optional
from fastapi import FastAPI, HTTPException, Security, Depends
from fastapi.middleware.cors import CORSMiddleware
from fastapi.security.api_key import APIKeyHeader
from pydantic import BaseModel
logging.basicConfig(level=logging.INFO, format="[%(asctime)s] %(levelname)s: %(message)s")
logger = logging.getLogger("chingubot-ai")
# ─────────────────────────────────────────────────────────────────────────────
# Config
# ─────────────────────────────────────────────────────────────────────────────
MODEL_NAME = os.getenv("EMBED_MODEL", "BAAI/bge-m3")
VECTOR_SIZE = int(os.getenv("VECTOR_SIZE", "1024"))
API_SECRET = os.getenv("HF_API_SECRET", "")
MAX_BATCH = int(os.getenv("MAX_BATCH", "64"))
# ─────────────────────────────────────────────────────────────────────────────
# Model loader
# ─────────────────────────────────────────────────────────────────────────────
_model = None
def get_model():
global _model
if _model is None:
raise HTTPException(status_code=503, detail="Model belum siap, coba lagi sebentar")
return _model
@asynccontextmanager
async def lifespan(app: FastAPI):
global _model
logger.info(f"Loading model: {MODEL_NAME}")
t0 = time.time()
try:
from sentence_transformers import SentenceTransformer
_model = SentenceTransformer(MODEL_NAME)
logger.info(f"βœ… Model loaded in {time.time()-t0:.1f}s | vector_size={VECTOR_SIZE}")
except Exception as e:
logger.error(f"❌ Model load failed: {e}")
_model = None
yield
logger.info("Shutting down...")
# ─────────────────────────────────────────────────────────────────────────────
# App
# ─────────────────────────────────────────────────────────────────────────────
app = FastAPI(title="ChinguBot AI Server", version="2.0.0", lifespan=lifespan)
app.add_middleware(
CORSMiddleware,
allow_origins=["*"],
allow_methods=["POST", "GET"],
allow_headers=["*"],
)
# ─────────────────────────────────────────────────────────────────────────────
# Auth
# ─────────────────────────────────────────────────────────────────────────────
api_key_header = APIKeyHeader(name="X-API-Key", auto_error=False)
def verify_key(key: Optional[str] = Security(api_key_header)):
if not API_SECRET:
return True
if key != API_SECRET:
raise HTTPException(status_code=401, detail="Invalid API key")
return True
# ─────────────────────────────────────────────────────────────────────────────
# Schemas
# ─────────────────────────────────────────────────────────────────────────────
class EmbedRequest(BaseModel):
text: str
class EmbedBatchRequest(BaseModel):
texts: List[str]
type: str = "query" # "query" | "passage"
class EmbedResponse(BaseModel):
vector: List[float]
size: int
model: str
class EmbedBatchResponse(BaseModel):
vectors: List[List[float]]
size: int
count: int
model: str
# ─────────────────────────────────────────────────────────────────────────────
# Endpoints
# ─────────────────────────────────────────────────────────────────────────────
@app.get("/ping")
def ping():
return {"status": "alive", "model": MODEL_NAME}
@app.get("/health")
def health():
return {
"status": "ok" if _model is not None else "loading",
"model": MODEL_NAME,
"vector_size": VECTOR_SIZE,
}
@app.post("/embed", response_model=EmbedResponse, dependencies=[Depends(verify_key)])
def embed_query(req: EmbedRequest):
"""Embed query (pesan user) β€” pakai prompt khusus retrieval."""
model = get_model()
if not req.text.strip():
raise HTTPException(status_code=400, detail="text kosong")
try:
vector = model.encode(
req.text.strip(),
normalize_embeddings=True,
).tolist()
return EmbedResponse(vector=vector, size=len(vector), model=MODEL_NAME)
except Exception as e:
logger.error(f"embed error: {e}")
raise HTTPException(status_code=500, detail=str(e))
@app.post("/embed-passage", response_model=EmbedResponse, dependencies=[Depends(verify_key)])
def embed_passage(req: EmbedRequest):
"""Embed passage/dokumen β€” untuk indexing intent utterances."""
model = get_model()
if not req.text.strip():
raise HTTPException(status_code=400, detail="text kosong")
try:
vector = model.encode(
req.text.strip(),
normalize_embeddings=True,
).tolist()
return EmbedResponse(vector=vector, size=len(vector), model=MODEL_NAME)
except Exception as e:
logger.error(f"embed-passage error: {e}")
raise HTTPException(status_code=500, detail=str(e))
@app.post("/embed-batch", response_model=EmbedBatchResponse, dependencies=[Depends(verify_key)])
def embed_batch(req: EmbedBatchRequest):
"""Batch embed β€” untuk seeding Qdrant atau embed history messages."""
model = get_model()
if not req.texts:
raise HTTPException(status_code=400, detail="texts kosong")
if len(req.texts) > MAX_BATCH:
raise HTTPException(status_code=400, detail=f"max {MAX_BATCH} texts per request")
cleaned = [t.strip() for t in req.texts if t.strip()]
if not cleaned:
raise HTTPException(status_code=400, detail="semua texts kosong")
try:
prompt_name = "query" if req.type == "query" else None
vecs = model.encode(
cleaned,
normalize_embeddings=True,
batch_size=32,
).tolist()
return EmbedBatchResponse(
vectors=vecs,
size=len(vecs[0]) if vecs else 0,
count=len(vecs),
model=MODEL_NAME,
)
except Exception as e:
logger.error(f"embed-batch error: {e}")
raise HTTPException(status_code=500, detail=str(e))