embeddinggemma_300m_fastapi / app_bk_260112.py
WildOjisan's picture
.
728f191
Raw History Blame Contribute Delete
7.7 kB
from contextlib import asynccontextmanager
from fastapi import FastAPI, Depends, HTTPException
from pydantic import BaseModel
import uvicorn
import asyncpg
from sentence_transformers import SentenceTransformer
import torch
from typing import List, Dict, Any, Union
from database_conn import connect_to_db, close_db_connection, get_db_connection
from test_router import router as test_router
import os # ํ™˜๊ฒฝ ๋ณ€์ˆ˜๋ฅผ ์ฝ๊ธฐ ์œ„ํ•ด os ๋ชจ๋“ˆ ์ถ”๊ฐ€
HF_TOKEN = os.getenv("HUGGING_FACE_HUB_TOKEN")
# Load the SentenceTransformer model once when the application starts
# This is a synchronous operation, which is fine for app startup.
try:
device = "cuda" if torch.cuda.is_available() else "cpu"
model = SentenceTransformer(
"google/embeddinggemma-300m",
device=device,
token=HF_TOKEN
)
# Set the model to evaluation mode
model.eval()
print("Model 'google/embeddinggemma-300m' loaded successfully.")
except Exception as e:
print(f"Error loading model: {e}")
# In a real application, you might want to raise an exception or handle this more gracefully
# For simplicity, we'll let the app potentially fail if the model can't load.
model = None # Keep model as None if loading failed
# 2. ๋ฐ์ดํ„ฐ ์œ ํšจ์„ฑ ๊ฒ€์‚ฌ๋ฅผ ์œ„ํ•œ Pydantic ๋ชจ๋ธ ์ •์˜
# ํด๋ผ์ด์–ธํŠธ๋กœ๋ถ€ํ„ฐ ๋ฐ›์„ ์š”์ฒญ(Request) ๋ฐ์ดํ„ฐ ๊ตฌ์กฐ๋ฅผ ์ •์˜ํ•ฉ๋‹ˆ๋‹ค.
class Item(BaseModel):
name: str
price: float
is_offer: bool | None = None
@asynccontextmanager
async def lifespan(app: FastAPI):
try:
await connect_to_db()
except Exception as e:
print(f"!!! [Startup] DB Connection FAILED: {e!r}")
# DB ์—ฐ๊ฒฐ ์‹คํŒจ ์‹œ์—๋„ ์„œ๋ฒ„๋ฅผ ๋„์šธ์ง€ ์—ฌ๋ถ€๋Š” ์„œ๋น„์Šค ์ •์ฑ…์— ๋”ฐ๋ผ ๊ฒฐ์ •
# --- ์„œ๋ฒ„ ์‹คํ–‰ ์‹œ์ž‘ (Yield) ---
yield
# --- ์„œ๋ฒ„ ์ข…๋ฃŒ ์‹œ (Shutdown Event) ---
await close_db_connection()
print(">>> [Shutdown] FastAPI Server graceful shutdown complete.")
app = FastAPI(
title="Gemma Embedding Service",
description="Implements text embedding generation via REST API.",
version="1.0.0",
lifespan=lifespan # <-- ์—ฌ๊ธฐ์—์„œ lifespan ํ•จ์ˆ˜๋ฅผ ๋“ฑ๋กํ•ฉ๋‹ˆ๋‹ค.
)
# --- API ์—”๋“œํฌ์ธํŠธ ์ •์˜ ---
# 3. ๋ฃจํŠธ ์—”๋“œํฌ์ธํŠธ (GET /)
@app.get("/")
def read_root():
result={"success":True,"data":None,"msg":""}
try:
result["data"]="ok"
return result
except Exception as e:
result["success"] = False
result["msg"]=f"server error. {e!r}"
return result
app.include_router(test_router, prefix="/api/test")
class MakeTextEmbedding(BaseModel):
"""
Input structure for the similarity endpoint.
"""
query: str
documents: List[str]
# ์ถœ๋ ฅ ํด๋ž˜์Šค: ๊ฒฐ๊ณผ๊ฐ€ ์„ฑ๊ณต ์—ฌ๋ถ€์™€ ์ž„๋ฒ ๋”ฉ ๋ฐ์ดํ„ฐ ๋ฆฌ์ŠคํŠธ๋ฅผ ํฌํ•จํ•˜๋„๋ก ์ •์˜
class EmbeddingOutput(BaseModel):
success: bool
msg: str
# ์ž„๋ฒ ๋”ฉ ๊ฒฐ๊ณผ๋Š” ์ค‘์ฒฉ ๋ฆฌ์ŠคํŠธ ํ˜•ํƒœ๊ฐ€ ๋ฉ๋‹ˆ๋‹ค. (๋ฌธ์„œ ์ˆ˜, ์ž„๋ฒ ๋”ฉ ์ฐจ์›)
data: Union[List[List[float]], None] = None
@app.post("/make_text_embedding", summary="Calculate semantic similarity and find the best match")
async def calculate_similarity(data: MakeTextEmbedding):
result={"success":True,"data":None,"msg":""}
try:
if model is None:
result["success"] = False
result["msg"]=f"Model not loaded. Service is unavailable."
return result
# The 'with torch.no_grad():' block is essential for efficient inference
with torch.no_grad():
# Encode the query (single vector)
document_embeddings = model.encode_document(data.documents)
embeddings_list = document_embeddings.tolist()
result["data"]=embeddings_list
result["msg"]=f"document_embeddings.shape:{document_embeddings.shape}"
return result
except Exception as e:
result["success"] = False
result["msg"]=f"server error. {e!r}"
return result
@app.post("/string_distance_compare", summary="Calculate semantic similarity and find the best match")
async def string_distance_compare(data: MakeTextEmbedding):
result={"success":True,"data":None,"msg":""}
try:
if model is None:
result["success"] = False
result["msg"]=f"Model not loaded. Service is unavailable."
return result
# The 'with torch.no_grad():' block is essential for efficient inference
with torch.no_grad():
# Encode the query (single vector)
query_embeddings = model.encode_query(data.query)
document_embeddings = model.encode_document(data.documents)
similarities = model.similarity(query_embeddings, document_embeddings)
result["data"]=similarities.tolist()
result["msg"]=f"query_embeddings.shape: {query_embeddings.shape}, document_embeddings.shape: {document_embeddings.shape}"
return result
except Exception as e:
result["success"] = False
result["msg"]=f"server error. {e!r}"
return result
# ----------------------------------------------------
# 2. API ์—”๋“œํฌ์ธํŠธ: SELECT NOW() (์Œฉ ์ฟผ๋ฆฌ ์‹คํ–‰)
# ----------------------------------------------------
@app.get("/time", response_model=Dict[str, Any])
async def get_db_time(
# get_db_connection์„ ํ†ตํ•ด asyncpg.Connection ๊ฐ์ฒด๋งŒ ์ฃผ์ž…๋ฐ›์Šต๋‹ˆ๋‹ค.
conn: asyncpg.Connection = Depends(get_db_connection)
):
"""
DB์— ์ ‘์†ํ•˜์—ฌ ํ˜„์žฌ ์‹œ๊ฐ„์„ ์กฐํšŒํ•˜๋Š” ์Œฉ SQL ์ฟผ๋ฆฌ(SELECT NOW())๋ฅผ ์‹คํ–‰ํ•ฉ๋‹ˆ๋‹ค.
"""
result = {"success": True, "data": None, "msg": ""}
try:
# **์—ฌ๊ธฐ์„œ ์Œฉ SQL ์ฟผ๋ฆฌ๋ฅผ ์ง์ ‘ ์ž‘์„ฑํ•˜๊ณ  ์‹คํ–‰ํ•ฉ๋‹ˆ๋‹ค.**
query = "SELECT NOW();"
# ์ฟผ๋ฆฌ ์‹คํ–‰ (fetchval()์€ ์ฟผ๋ฆฌ ๊ฒฐ๊ณผ์˜ ์ฒซ ํ–‰, ์ฒซ ์—ด์˜ ๊ฐ’๋งŒ ๋ฐ˜ํ™˜ํ•ฉ๋‹ˆ๋‹ค.)
# asyncpg๊ฐ€ DB ์‹œ๊ฐ„์„ Python datetime ๊ฐ์ฒด๋กœ ๋ณ€ํ™˜ํ•ด ์ค๋‹ˆ๋‹ค.
records = await conn.fetch(query)
data_list_of_dicts: List[Dict[str, Any]] = [
# dict(record)๋ฅผ ์‚ฌ์šฉํ•˜์—ฌ Record ๊ฐ์ฒด๋ฅผ ์ผ๋ฐ˜ ํŒŒ์ด์ฌ ๋”•์…”๋„ˆ๋ฆฌ๋กœ ๋ณ€ํ™˜
dict(record) for record in records
]
# ๊ฒฐ๊ณผ๋ฅผ ๋ฌธ์ž์—ด๋กœ ๋ณ€ํ™˜ํ•˜์—ฌ JSON ์‘๋‹ต์— ๋‹ด์Šต๋‹ˆ๋‹ค.
result["data"] = data_list_of_dicts[0]
except Exception as e:
# ์ฟผ๋ฆฌ ์‹คํ–‰ ์ค‘ ์˜ค๋ฅ˜๊ฐ€ ๋ฐœ์ƒํ•˜๋ฉด 500 ์—๋Ÿฌ๋ฅผ ๋ฐ˜ํ™˜
result["success"] = False
result["msg"] = f"Database query error: {e!r}"
return result
# 4. ์ฟผ๋ฆฌ ํŒŒ๋ผ๋ฏธํ„ฐ๋ฅผ ๋ฐ›๋Š” ์—”๋“œํฌ์ธํŠธ (GET /items/{item_id})
@app.get("/items")
def read_item(q: str | None = None):
"""
ํŠน์ • ID์˜ ์•„์ดํ…œ ์ •๋ณด๋ฅผ ๊ฐ€์ ธ์˜ต๋‹ˆ๋‹ค.
:param item_id: ์•„์ดํ…œ์˜ ๊ณ ์œ  ID (๊ฒฝ๋กœ ๋งค๊ฐœ๋ณ€์ˆ˜)
:param q: ์„ ํƒ์ ์ธ ๊ฒ€์ƒ‰ ๋ฌธ์ž์—ด (์ฟผ๋ฆฌ ๋งค๊ฐœ๋ณ€์ˆ˜)
"""
return {"q": q, "description": "This is a query test."}
# 5. ์š”์ฒญ ๋ณธ๋ฌธ(Body)์„ ๋ฐ›๋Š” ์—”๋“œํฌ์ธํŠธ (POST /items/)
@app.post("/items/")
def create_item(item: Item):
"""
์ƒˆ๋กœ์šด ์•„์ดํ…œ์„ ์ƒ์„ฑํ•˜๊ณ  ์ •๋ณด๋ฅผ ๋ฐ˜ํ™˜ํ•ฉ๋‹ˆ๋‹ค.
:param item: Item Pydantic ๋ชจ๋ธ์— ์ •์˜๋œ ๋ฐ์ดํ„ฐ ๊ตฌ์กฐ
"""
# ์‹ค์ œ ๋ฐ์ดํ„ฐ๋ฒ ์ด์Šค์— ์ €์žฅํ•˜๋Š” ๋Œ€์‹ , ๊ฐ„๋‹จํ•œ ์ฒ˜๋ฆฌ๋ฅผ ์ˆ˜ํ–‰
if item.price > 100.0:
item.name = f"Premium {item.name}"
return {"message": "Item created successfully", "item_data": item}
if __name__ == "__main__":
# --reload ์˜ต์…˜์„ ์ถ”๊ฐ€ํ•˜์—ฌ ์ฝ”๋“œ๊ฐ€ ๋ณ€๊ฒฝ๋  ๋•Œ๋งˆ๋‹ค ์ž๋™ ์žฌ์‹œ์ž‘๋˜๊ฒŒ ์„ค์ •ํ•ฉ๋‹ˆ๋‹ค.
uvicorn.run("app:app", host="0.0.0.0", port=8000, reload=True)