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)