774 lines
24 KiB
Python
774 lines
24 KiB
Python
|
|
"""Phase 7 FastAPI application: LLM End-to-End Integration.
|
||
|
|
|
||
|
|
Features:
|
||
|
|
- Direct LLM integration (OpenAI, Anthropic, Local)
|
||
|
|
- Response streaming (Server-Sent Events)
|
||
|
|
- Redis caching with TTL
|
||
|
|
- RAG + LLM unified pipeline
|
||
|
|
- Multiple LLM provider support
|
||
|
|
"""
|
||
|
|
|
||
|
|
import asyncio
|
||
|
|
import hashlib
|
||
|
|
import json
|
||
|
|
import logging
|
||
|
|
import time
|
||
|
|
from typing import Optional, Dict, Any, AsyncGenerator
|
||
|
|
from datetime import datetime, timedelta
|
||
|
|
|
||
|
|
from fastapi import FastAPI, APIRouter, HTTPException, Query, Request, Depends
|
||
|
|
from fastapi.responses import StreamingResponse, JSONResponse
|
||
|
|
from pydantic import BaseModel, Field
|
||
|
|
|
||
|
|
try:
|
||
|
|
import redis.asyncio as redis
|
||
|
|
REDIS_AVAILABLE = True
|
||
|
|
except ImportError:
|
||
|
|
REDIS_AVAILABLE = False
|
||
|
|
|
||
|
|
from ont_platform.core.graph.neo4j_adapter import Neo4jAdapter, Neo4jConfig
|
||
|
|
from ont_platform.core.graph.entity_resolver import EntityResolver
|
||
|
|
from ont_platform.core.graph.subgraph_retriever import SubgraphRetriever
|
||
|
|
from ont_platform.core.graph.pattern_matcher import PatternMatcher
|
||
|
|
from ont_platform.core.graph.graph_analytics import GraphAnalytics
|
||
|
|
from ont_platform.llm.llm_integration import LLMManager, LLMConfig, LLMProvider
|
||
|
|
|
||
|
|
logger = logging.getLogger(__name__)
|
||
|
|
|
||
|
|
# ============================================================================
|
||
|
|
# Request/Response Models
|
||
|
|
# ============================================================================
|
||
|
|
|
||
|
|
|
||
|
|
class AskRequest(BaseModel):
|
||
|
|
"""LLM query request."""
|
||
|
|
query: str = Field(..., description="사용자 질문")
|
||
|
|
context_hops: int = Field(2, description="그래프 컨텍스트 깊이")
|
||
|
|
use_cache: bool = Field(True, description="캐시 사용 여부")
|
||
|
|
temperature: Optional[float] = Field(None, description="LLM 온도 (0~1)")
|
||
|
|
max_tokens: Optional[int] = Field(None, description="최대 토큰 수")
|
||
|
|
|
||
|
|
|
||
|
|
class AskResponse(BaseModel):
|
||
|
|
"""LLM query response."""
|
||
|
|
query: str
|
||
|
|
answer: str
|
||
|
|
context_size: int
|
||
|
|
relevant_entities: list[str]
|
||
|
|
latency_ms: float
|
||
|
|
cached: bool = False
|
||
|
|
model: str
|
||
|
|
provider: str
|
||
|
|
|
||
|
|
|
||
|
|
class StreamingAskRequest(BaseModel):
|
||
|
|
"""Streaming LLM query request."""
|
||
|
|
query: str
|
||
|
|
context_hops: int = 2
|
||
|
|
temperature: Optional[float] = None
|
||
|
|
max_tokens: Optional[int] = None
|
||
|
|
|
||
|
|
|
||
|
|
class RAGMetadata(BaseModel):
|
||
|
|
"""RAG metadata."""
|
||
|
|
query: str
|
||
|
|
context_nodes: int
|
||
|
|
relevant_entities: list[str]
|
||
|
|
extraction_time_ms: float
|
||
|
|
llm_provider: str
|
||
|
|
llm_model: str
|
||
|
|
|
||
|
|
|
||
|
|
# ============================================================================
|
||
|
|
# Global Instances
|
||
|
|
# ============================================================================
|
||
|
|
|
||
|
|
_neo4j_adapter: Optional[Neo4jAdapter] = None
|
||
|
|
_entity_resolver: Optional[EntityResolver] = None
|
||
|
|
_subgraph_retriever: Optional[SubgraphRetriever] = None
|
||
|
|
_pattern_matcher: Optional[PatternMatcher] = None
|
||
|
|
_graph_analytics: Optional[GraphAnalytics] = None
|
||
|
|
_llm_manager: Optional[LLMManager] = None
|
||
|
|
_redis_client: Optional[redis.Redis] = None
|
||
|
|
|
||
|
|
# ============================================================================
|
||
|
|
# FastAPI App
|
||
|
|
# ============================================================================
|
||
|
|
|
||
|
|
app = FastAPI(
|
||
|
|
title="Ontology Platform - Phase 7 LLM Integration",
|
||
|
|
description="LLM End-to-End Integration with Streaming & Caching",
|
||
|
|
version="0.7.0",
|
||
|
|
)
|
||
|
|
|
||
|
|
llm_router = APIRouter(prefix="/api/v1/llm", tags=["llm"])
|
||
|
|
|
||
|
|
|
||
|
|
# ============================================================================
|
||
|
|
# Initialization Functions
|
||
|
|
# ============================================================================
|
||
|
|
|
||
|
|
|
||
|
|
async def get_neo4j_adapter() -> Neo4jAdapter:
|
||
|
|
"""Get or create Neo4j adapter instance."""
|
||
|
|
global _neo4j_adapter
|
||
|
|
if _neo4j_adapter is None:
|
||
|
|
config = Neo4jConfig(
|
||
|
|
uri="bolt://localhost:7687",
|
||
|
|
username="neo4j",
|
||
|
|
password="ontology123",
|
||
|
|
)
|
||
|
|
_neo4j_adapter = Neo4jAdapter(config)
|
||
|
|
if not await _neo4j_adapter.connect():
|
||
|
|
logger.warning("Neo4j not available")
|
||
|
|
else:
|
||
|
|
try:
|
||
|
|
await _neo4j_adapter.initialize_embedder()
|
||
|
|
except Exception as e:
|
||
|
|
logger.warning(f"Failed to initialize embedder: {e}")
|
||
|
|
return _neo4j_adapter
|
||
|
|
|
||
|
|
|
||
|
|
async def get_components():
|
||
|
|
"""Initialize all graph components."""
|
||
|
|
global (
|
||
|
|
_entity_resolver,
|
||
|
|
_subgraph_retriever,
|
||
|
|
_pattern_matcher,
|
||
|
|
_graph_analytics,
|
||
|
|
)
|
||
|
|
|
||
|
|
adapter = await get_neo4j_adapter()
|
||
|
|
|
||
|
|
if _entity_resolver is None:
|
||
|
|
_entity_resolver = EntityResolver()
|
||
|
|
await _entity_resolver.initialize_embedder()
|
||
|
|
|
||
|
|
if _subgraph_retriever is None:
|
||
|
|
_subgraph_retriever = SubgraphRetriever(adapter)
|
||
|
|
|
||
|
|
if _pattern_matcher is None:
|
||
|
|
_pattern_matcher = PatternMatcher(adapter)
|
||
|
|
|
||
|
|
if _graph_analytics is None:
|
||
|
|
_graph_analytics = GraphAnalytics(adapter)
|
||
|
|
|
||
|
|
return {
|
||
|
|
"adapter": adapter,
|
||
|
|
"resolver": _entity_resolver,
|
||
|
|
"retriever": _subgraph_retriever,
|
||
|
|
"matcher": _pattern_matcher,
|
||
|
|
"analytics": _graph_analytics,
|
||
|
|
}
|
||
|
|
|
||
|
|
|
||
|
|
async def get_llm_manager() -> LLMManager:
|
||
|
|
"""Get or create LLM manager instance."""
|
||
|
|
global _llm_manager
|
||
|
|
if _llm_manager is None:
|
||
|
|
# Default to OpenAI, but can be overridden via environment
|
||
|
|
config = LLMConfig(
|
||
|
|
provider=LLMProvider.OPENAI,
|
||
|
|
api_key=None, # Will use OPENAI_API_KEY env
|
||
|
|
model="gpt-4",
|
||
|
|
temperature=0.7,
|
||
|
|
max_tokens=500,
|
||
|
|
)
|
||
|
|
_llm_manager = LLMManager(config)
|
||
|
|
return _llm_manager
|
||
|
|
|
||
|
|
|
||
|
|
async def get_redis_client() -> Optional[redis.Redis]:
|
||
|
|
"""Get or create Redis client instance."""
|
||
|
|
global _redis_client
|
||
|
|
if not REDIS_AVAILABLE:
|
||
|
|
return None
|
||
|
|
|
||
|
|
if _redis_client is None:
|
||
|
|
try:
|
||
|
|
_redis_client = await redis.from_url(
|
||
|
|
"redis://localhost:6379",
|
||
|
|
decode_responses=True
|
||
|
|
)
|
||
|
|
await _redis_client.ping()
|
||
|
|
logger.info("Redis connected successfully")
|
||
|
|
except Exception as e:
|
||
|
|
logger.warning(f"Redis not available: {e}")
|
||
|
|
_redis_client = None
|
||
|
|
return _redis_client
|
||
|
|
|
||
|
|
|
||
|
|
# ============================================================================
|
||
|
|
# Cache Utilities
|
||
|
|
# ============================================================================
|
||
|
|
|
||
|
|
|
||
|
|
def _generate_cache_key(query: str, context_hops: int) -> str:
|
||
|
|
"""Generate cache key from query and context."""
|
||
|
|
key_data = f"{query}:{context_hops}"
|
||
|
|
key_hash = hashlib.sha256(key_data.encode()).hexdigest()[:16]
|
||
|
|
return f"phase7:rag:{key_hash}"
|
||
|
|
|
||
|
|
|
||
|
|
async def _get_cached_response(
|
||
|
|
redis_client: Optional[redis.Redis],
|
||
|
|
cache_key: str
|
||
|
|
) -> Optional[Dict[str, Any]]:
|
||
|
|
"""Retrieve cached response from Redis."""
|
||
|
|
if not redis_client:
|
||
|
|
return None
|
||
|
|
|
||
|
|
try:
|
||
|
|
cached = await redis_client.get(cache_key)
|
||
|
|
if cached:
|
||
|
|
logger.info(f"Cache hit: {cache_key}")
|
||
|
|
return json.loads(cached)
|
||
|
|
except Exception as e:
|
||
|
|
logger.warning(f"Cache retrieval failed: {e}")
|
||
|
|
|
||
|
|
return None
|
||
|
|
|
||
|
|
|
||
|
|
async def _cache_response(
|
||
|
|
redis_client: Optional[redis.Redis],
|
||
|
|
cache_key: str,
|
||
|
|
response: Dict[str, Any],
|
||
|
|
ttl_hours: int = 1
|
||
|
|
) -> bool:
|
||
|
|
"""Cache response in Redis."""
|
||
|
|
if not redis_client:
|
||
|
|
return False
|
||
|
|
|
||
|
|
try:
|
||
|
|
ttl_seconds = ttl_hours * 3600
|
||
|
|
await redis_client.setex(
|
||
|
|
cache_key,
|
||
|
|
ttl_seconds,
|
||
|
|
json.dumps(response, default=str)
|
||
|
|
)
|
||
|
|
logger.info(f"Cached response: {cache_key} (TTL: {ttl_hours}h)")
|
||
|
|
return True
|
||
|
|
except Exception as e:
|
||
|
|
logger.warning(f"Cache storage failed: {e}")
|
||
|
|
return False
|
||
|
|
|
||
|
|
|
||
|
|
# ============================================================================
|
||
|
|
# RAG Context Extraction
|
||
|
|
# ============================================================================
|
||
|
|
|
||
|
|
|
||
|
|
async def extract_rag_context(
|
||
|
|
query: str,
|
||
|
|
context_hops: int = 2,
|
||
|
|
max_entities: int = 100,
|
||
|
|
) -> Dict[str, Any]:
|
||
|
|
"""Extract RAG context from knowledge graph."""
|
||
|
|
start_time = time.time()
|
||
|
|
components = await get_components()
|
||
|
|
|
||
|
|
try:
|
||
|
|
# 1. Find relevant entities by semantic similarity
|
||
|
|
# Using entity resolver's embedding capability
|
||
|
|
retriever = components["retriever"]
|
||
|
|
|
||
|
|
# For now, retrieve a default context
|
||
|
|
# In production, would search by query semantic similarity
|
||
|
|
context_data = {
|
||
|
|
"query": query,
|
||
|
|
"nodes": [],
|
||
|
|
"edges": [],
|
||
|
|
"relevant_entities": [],
|
||
|
|
}
|
||
|
|
|
||
|
|
# Try to get context from first few entities as example
|
||
|
|
try:
|
||
|
|
# Get graph statistics to find some entities
|
||
|
|
analytics = components["analytics"]
|
||
|
|
stats = await analytics.get_graph_statistics()
|
||
|
|
|
||
|
|
if stats.get("total_nodes", 0) > 0:
|
||
|
|
# Get influential entities as relevant context
|
||
|
|
influential = await analytics.find_influential_entities(top_n=5)
|
||
|
|
context_data["relevant_entities"] = [
|
||
|
|
e.get("label", f"Entity_{e.get('entity_id')}")
|
||
|
|
for e in influential
|
||
|
|
]
|
||
|
|
context_data["nodes"] = influential[:max_entities]
|
||
|
|
except Exception as e:
|
||
|
|
logger.warning(f"Failed to extract context: {e}")
|
||
|
|
context_data["relevant_entities"] = []
|
||
|
|
|
||
|
|
extraction_time = (time.time() - start_time) * 1000
|
||
|
|
context_data["extraction_time_ms"] = extraction_time
|
||
|
|
|
||
|
|
return context_data
|
||
|
|
except Exception as e:
|
||
|
|
logger.error(f"RAG context extraction failed: {e}")
|
||
|
|
raise
|
||
|
|
|
||
|
|
|
||
|
|
# ============================================================================
|
||
|
|
# RAG Prompt Building
|
||
|
|
# ============================================================================
|
||
|
|
|
||
|
|
|
||
|
|
def _build_rag_prompt_for_llm(
|
||
|
|
query: str,
|
||
|
|
context: Dict[str, Any]
|
||
|
|
) -> str:
|
||
|
|
"""Build structured prompt with RAG context for LLM."""
|
||
|
|
relevant_entities = context.get("relevant_entities", [])
|
||
|
|
nodes = context.get("nodes", [])
|
||
|
|
|
||
|
|
# Build context section
|
||
|
|
context_str = ""
|
||
|
|
if relevant_entities:
|
||
|
|
context_str += "관련 엔티티:\n"
|
||
|
|
for entity in relevant_entities[:10]:
|
||
|
|
if isinstance(entity, dict):
|
||
|
|
label = entity.get("label", "Unknown")
|
||
|
|
entity_type = entity.get("type", "Unknown")
|
||
|
|
else:
|
||
|
|
label = str(entity)
|
||
|
|
entity_type = "Unknown"
|
||
|
|
context_str += f"- {label} ({entity_type})\n"
|
||
|
|
|
||
|
|
if nodes:
|
||
|
|
context_str += "\n그래프 정보:\n"
|
||
|
|
for node in nodes[:5]:
|
||
|
|
if isinstance(node, dict):
|
||
|
|
label = node.get("label", "Unknown")
|
||
|
|
context_str += f"- {label}\n"
|
||
|
|
|
||
|
|
# Build system prompt with context
|
||
|
|
prompt = f"""당신은 지식 그래프 기반 질문 답변 어시스턴트입니다.
|
||
|
|
다음 지식 그래프 정보를 참고하여 질문에 답변해주세요.
|
||
|
|
|
||
|
|
=== 지식 그래프 컨텍스트 ===
|
||
|
|
{context_str if context_str else "컨텍스트 없음"}
|
||
|
|
|
||
|
|
=== 사용자 질문 ===
|
||
|
|
{query}
|
||
|
|
|
||
|
|
위의 지식 그래프 정보를 바탕으로 명확하고 정확한 답변을 제공해주세요."""
|
||
|
|
|
||
|
|
return prompt
|
||
|
|
|
||
|
|
|
||
|
|
# ============================================================================
|
||
|
|
# LLM Endpoints
|
||
|
|
# ============================================================================
|
||
|
|
|
||
|
|
|
||
|
|
@llm_router.post("/ask", response_model=AskResponse)
|
||
|
|
async def ask_llm(request: AskRequest) -> AskResponse:
|
||
|
|
"""
|
||
|
|
LLM에 질문을 하고 캐시된 응답을 반환합니다.
|
||
|
|
|
||
|
|
- RAG 컨텍스트 자동 추출
|
||
|
|
- Redis 캐싱 (기본 1시간 TTL)
|
||
|
|
- 단일 응답 반환
|
||
|
|
"""
|
||
|
|
start_time = time.time()
|
||
|
|
|
||
|
|
# 캐시 확인
|
||
|
|
redis_client = await get_redis_client()
|
||
|
|
cache_key = _generate_cache_key(request.query, request.context_hops)
|
||
|
|
|
||
|
|
if request.use_cache:
|
||
|
|
cached = await _get_cached_response(redis_client, cache_key)
|
||
|
|
if cached:
|
||
|
|
cached["cached"] = True
|
||
|
|
cached["latency_ms"] = (time.time() - start_time) * 1000
|
||
|
|
return AskResponse(**cached)
|
||
|
|
|
||
|
|
try:
|
||
|
|
# RAG 컨텍스트 추출
|
||
|
|
context_start = time.time()
|
||
|
|
context = await extract_rag_context(
|
||
|
|
request.query,
|
||
|
|
context_hops=request.context_hops
|
||
|
|
)
|
||
|
|
context_time = (time.time() - context_start) * 1000
|
||
|
|
|
||
|
|
# 프롬프트 생성
|
||
|
|
prompt = _build_rag_prompt_for_llm(request.query, context)
|
||
|
|
|
||
|
|
# LLM 호출
|
||
|
|
llm_manager = await get_llm_manager()
|
||
|
|
llm_start = time.time()
|
||
|
|
|
||
|
|
# LLM 설정 업데이트 (요청으로부터)
|
||
|
|
if request.temperature is not None:
|
||
|
|
llm_manager.config.temperature = request.temperature
|
||
|
|
if request.max_tokens is not None:
|
||
|
|
llm_manager.config.max_tokens = request.max_tokens
|
||
|
|
|
||
|
|
answer = await llm_manager.generate(prompt, stream=False)
|
||
|
|
llm_time = (time.time() - llm_start) * 1000
|
||
|
|
|
||
|
|
# 응답 생성
|
||
|
|
response_data = {
|
||
|
|
"query": request.query,
|
||
|
|
"answer": answer,
|
||
|
|
"context_size": len(context.get("nodes", [])),
|
||
|
|
"relevant_entities": context.get("relevant_entities", []),
|
||
|
|
"latency_ms": (time.time() - start_time) * 1000,
|
||
|
|
"cached": False,
|
||
|
|
"model": llm_manager.config.model,
|
||
|
|
"provider": llm_manager.config.provider.value,
|
||
|
|
}
|
||
|
|
|
||
|
|
# 응답 캐시
|
||
|
|
if request.use_cache:
|
||
|
|
await _cache_response(redis_client, cache_key, response_data)
|
||
|
|
|
||
|
|
logger.info(
|
||
|
|
f"LLM query completed. "
|
||
|
|
f"Context: {context_time:.1f}ms, "
|
||
|
|
f"LLM: {llm_time:.1f}ms, "
|
||
|
|
f"Total: {response_data['latency_ms']:.1f}ms"
|
||
|
|
)
|
||
|
|
|
||
|
|
return AskResponse(**response_data)
|
||
|
|
|
||
|
|
except Exception as e:
|
||
|
|
logger.error(f"LLM query failed: {e}")
|
||
|
|
raise HTTPException(status_code=500, detail=str(e))
|
||
|
|
|
||
|
|
|
||
|
|
@llm_router.post("/ask/stream")
|
||
|
|
async def ask_llm_stream(request: StreamingAskRequest):
|
||
|
|
"""
|
||
|
|
LLM에 질문을 하고 스트리밍 응답을 반환합니다.
|
||
|
|
|
||
|
|
- Server-Sent Events (SSE) 기반 스트리밍
|
||
|
|
- 실시간 토큰 전달
|
||
|
|
- 메타데이터 포함
|
||
|
|
"""
|
||
|
|
|
||
|
|
async def stream_generator() -> AsyncGenerator[str, None]:
|
||
|
|
"""Stream LLM response tokens."""
|
||
|
|
try:
|
||
|
|
# RAG 컨텍스트 추출
|
||
|
|
context = await extract_rag_context(
|
||
|
|
request.query,
|
||
|
|
context_hops=request.context_hops
|
||
|
|
)
|
||
|
|
|
||
|
|
# 메타데이터 전송
|
||
|
|
metadata = {
|
||
|
|
"type": "metadata",
|
||
|
|
"query": request.query,
|
||
|
|
"context_nodes": len(context.get("nodes", [])),
|
||
|
|
"relevant_entities": context.get("relevant_entities", []),
|
||
|
|
"extraction_time_ms": context.get("extraction_time_ms", 0),
|
||
|
|
}
|
||
|
|
yield f"data: {json.dumps(metadata)}\n\n"
|
||
|
|
|
||
|
|
# 프롬프트 생성
|
||
|
|
prompt = _build_rag_prompt_for_llm(request.query, context)
|
||
|
|
|
||
|
|
# LLM 스트리밍 호출
|
||
|
|
llm_manager = await get_llm_manager()
|
||
|
|
|
||
|
|
if request.temperature is not None:
|
||
|
|
llm_manager.config.temperature = request.temperature
|
||
|
|
if request.max_tokens is not None:
|
||
|
|
llm_manager.config.max_tokens = request.max_tokens
|
||
|
|
|
||
|
|
# 토큰 스트리밍
|
||
|
|
token_count = 0
|
||
|
|
async for token in llm_manager.generate_stream(prompt):
|
||
|
|
token_data = {
|
||
|
|
"type": "token",
|
||
|
|
"content": token,
|
||
|
|
"token_index": token_count,
|
||
|
|
}
|
||
|
|
yield f"data: {json.dumps(token_data)}\n\n"
|
||
|
|
token_count += 1
|
||
|
|
|
||
|
|
# 완료 신호
|
||
|
|
completion = {
|
||
|
|
"type": "complete",
|
||
|
|
"total_tokens": token_count,
|
||
|
|
"timestamp": datetime.now().isoformat(),
|
||
|
|
}
|
||
|
|
yield f"data: {json.dumps(completion)}\n\n"
|
||
|
|
|
||
|
|
except Exception as e:
|
||
|
|
error_data = {
|
||
|
|
"type": "error",
|
||
|
|
"message": str(e),
|
||
|
|
}
|
||
|
|
yield f"data: {json.dumps(error_data)}\n\n"
|
||
|
|
logger.error(f"Streaming error: {e}")
|
||
|
|
|
||
|
|
return StreamingResponse(
|
||
|
|
stream_generator(),
|
||
|
|
media_type="text/event-stream",
|
||
|
|
headers={
|
||
|
|
"Cache-Control": "no-cache",
|
||
|
|
"X-Accel-Buffering": "no",
|
||
|
|
}
|
||
|
|
)
|
||
|
|
|
||
|
|
|
||
|
|
@llm_router.post("/ask/metadata")
|
||
|
|
async def get_rag_metadata(request: AskRequest) -> RAGMetadata:
|
||
|
|
"""
|
||
|
|
RAG 추출 메타데이터만 반환 (LLM 호출 없음).
|
||
|
|
|
||
|
|
- 컨텍스트 추출 시간만 측정
|
||
|
|
- 응답 최소화 (메타데이터만)
|
||
|
|
"""
|
||
|
|
try:
|
||
|
|
context_start = time.time()
|
||
|
|
context = await extract_rag_context(
|
||
|
|
request.query,
|
||
|
|
context_hops=request.context_hops
|
||
|
|
)
|
||
|
|
extraction_time = (time.time() - context_start) * 1000
|
||
|
|
|
||
|
|
llm_manager = await get_llm_manager()
|
||
|
|
|
||
|
|
return RAGMetadata(
|
||
|
|
query=request.query,
|
||
|
|
context_nodes=len(context.get("nodes", [])),
|
||
|
|
relevant_entities=context.get("relevant_entities", []),
|
||
|
|
extraction_time_ms=extraction_time,
|
||
|
|
llm_provider=llm_manager.config.provider.value,
|
||
|
|
llm_model=llm_manager.config.model,
|
||
|
|
)
|
||
|
|
except Exception as e:
|
||
|
|
logger.error(f"Metadata retrieval failed: {e}")
|
||
|
|
raise HTTPException(status_code=500, detail=str(e))
|
||
|
|
|
||
|
|
|
||
|
|
@llm_router.post("/configure")
|
||
|
|
async def configure_llm(
|
||
|
|
provider: str = Query(..., description="LLM Provider: openai, anthropic, local"),
|
||
|
|
model: str = Query(..., description="Model name"),
|
||
|
|
api_key: Optional[str] = Query(None, description="API key (optional)"),
|
||
|
|
temperature: float = Query(0.7, ge=0.0, le=2.0),
|
||
|
|
max_tokens: int = Query(500, ge=1, le=4000),
|
||
|
|
) -> Dict[str, Any]:
|
||
|
|
"""
|
||
|
|
LLM 설정 변경.
|
||
|
|
|
||
|
|
- Provider 변경 (OpenAI, Anthropic, Local)
|
||
|
|
- 모델 선택
|
||
|
|
- 온도/토큰 조정
|
||
|
|
"""
|
||
|
|
global _llm_manager
|
||
|
|
|
||
|
|
try:
|
||
|
|
provider_enum = LLMProvider[provider.upper()]
|
||
|
|
except KeyError:
|
||
|
|
raise HTTPException(
|
||
|
|
status_code=400,
|
||
|
|
detail=f"Unknown provider: {provider}. "
|
||
|
|
f"Choose from: {[p.value for p in LLMProvider]}"
|
||
|
|
)
|
||
|
|
|
||
|
|
try:
|
||
|
|
config = LLMConfig(
|
||
|
|
provider=provider_enum,
|
||
|
|
api_key=api_key,
|
||
|
|
model=model,
|
||
|
|
temperature=temperature,
|
||
|
|
max_tokens=max_tokens,
|
||
|
|
base_url="http://localhost:1234/v1" if provider_enum == LLMProvider.LOCAL else None,
|
||
|
|
)
|
||
|
|
_llm_manager = LLMManager(config)
|
||
|
|
|
||
|
|
return {
|
||
|
|
"status": "configured",
|
||
|
|
"provider": provider,
|
||
|
|
"model": model,
|
||
|
|
"temperature": temperature,
|
||
|
|
"max_tokens": max_tokens,
|
||
|
|
}
|
||
|
|
except Exception as e:
|
||
|
|
logger.error(f"LLM configuration failed: {e}")
|
||
|
|
raise HTTPException(status_code=500, detail=str(e))
|
||
|
|
|
||
|
|
|
||
|
|
@llm_router.get("/info")
|
||
|
|
async def get_llm_info() -> Dict[str, Any]:
|
||
|
|
"""Get current LLM configuration and status."""
|
||
|
|
try:
|
||
|
|
llm_manager = await get_llm_manager()
|
||
|
|
redis_client = await get_redis_client()
|
||
|
|
|
||
|
|
return {
|
||
|
|
"llm_provider": llm_manager.config.provider.value,
|
||
|
|
"llm_model": llm_manager.config.model,
|
||
|
|
"temperature": llm_manager.config.temperature,
|
||
|
|
"max_tokens": llm_manager.config.max_tokens,
|
||
|
|
"redis_available": redis_client is not None,
|
||
|
|
"timestamp": datetime.now().isoformat(),
|
||
|
|
}
|
||
|
|
except Exception as e:
|
||
|
|
logger.error(f"Failed to get LLM info: {e}")
|
||
|
|
raise HTTPException(status_code=500, detail=str(e))
|
||
|
|
|
||
|
|
|
||
|
|
# ============================================================================
|
||
|
|
# Cache Management
|
||
|
|
# ============================================================================
|
||
|
|
|
||
|
|
|
||
|
|
@llm_router.delete("/cache")
|
||
|
|
async def clear_cache() -> Dict[str, str]:
|
||
|
|
"""모든 RAG 캐시 삭제."""
|
||
|
|
redis_client = await get_redis_client()
|
||
|
|
if not redis_client:
|
||
|
|
return {"status": "redis_unavailable"}
|
||
|
|
|
||
|
|
try:
|
||
|
|
cursor = 0
|
||
|
|
deleted = 0
|
||
|
|
|
||
|
|
while True:
|
||
|
|
cursor, keys = await redis_client.scan(
|
||
|
|
cursor,
|
||
|
|
match="phase7:rag:*",
|
||
|
|
count=100
|
||
|
|
)
|
||
|
|
|
||
|
|
if keys:
|
||
|
|
await redis_client.delete(*keys)
|
||
|
|
deleted += len(keys)
|
||
|
|
|
||
|
|
if cursor == 0:
|
||
|
|
break
|
||
|
|
|
||
|
|
return {
|
||
|
|
"status": "success",
|
||
|
|
"deleted_keys": str(deleted),
|
||
|
|
}
|
||
|
|
except Exception as e:
|
||
|
|
logger.error(f"Cache clearing failed: {e}")
|
||
|
|
raise HTTPException(status_code=500, detail=str(e))
|
||
|
|
|
||
|
|
|
||
|
|
@llm_router.get("/cache/info")
|
||
|
|
async def get_cache_info() -> Dict[str, Any]:
|
||
|
|
"""캐시 통계."""
|
||
|
|
redis_client = await get_redis_client()
|
||
|
|
if not redis_client:
|
||
|
|
return {"redis_available": False}
|
||
|
|
|
||
|
|
try:
|
||
|
|
info = await redis_client.info()
|
||
|
|
cursor = 0
|
||
|
|
cache_keys = 0
|
||
|
|
|
||
|
|
while True:
|
||
|
|
cursor, keys = await redis_client.scan(
|
||
|
|
cursor,
|
||
|
|
match="phase7:rag:*",
|
||
|
|
count=100
|
||
|
|
)
|
||
|
|
cache_keys += len(keys)
|
||
|
|
|
||
|
|
if cursor == 0:
|
||
|
|
break
|
||
|
|
|
||
|
|
return {
|
||
|
|
"redis_available": True,
|
||
|
|
"used_memory_mb": info.get("used_memory", 0) / (1024 * 1024),
|
||
|
|
"cache_keys": cache_keys,
|
||
|
|
"redis_version": info.get("redis_version", "unknown"),
|
||
|
|
}
|
||
|
|
except Exception as e:
|
||
|
|
logger.error(f"Cache info retrieval failed: {e}")
|
||
|
|
return {"redis_available": False, "error": str(e)}
|
||
|
|
|
||
|
|
|
||
|
|
# ============================================================================
|
||
|
|
# Health & Info Endpoints
|
||
|
|
# ============================================================================
|
||
|
|
|
||
|
|
|
||
|
|
@app.get("/health")
|
||
|
|
async def health_check() -> Dict[str, Any]:
|
||
|
|
"""헬스 체크."""
|
||
|
|
try:
|
||
|
|
adapter = await get_neo4j_adapter()
|
||
|
|
neo4j_ok = adapter is not None and adapter.driver is not None
|
||
|
|
|
||
|
|
redis_client = await get_redis_client()
|
||
|
|
redis_ok = redis_client is not None
|
||
|
|
|
||
|
|
llm_manager = await get_llm_manager()
|
||
|
|
|
||
|
|
return {
|
||
|
|
"status": "healthy",
|
||
|
|
"version": "0.7.0",
|
||
|
|
"neo4j": "connected" if neo4j_ok else "disconnected",
|
||
|
|
"redis": "available" if redis_ok else "unavailable",
|
||
|
|
"llm_provider": llm_manager.config.provider.value,
|
||
|
|
"timestamp": datetime.now().isoformat(),
|
||
|
|
}
|
||
|
|
except Exception as e:
|
||
|
|
logger.error(f"Health check failed: {e}")
|
||
|
|
return {
|
||
|
|
"status": "unhealthy",
|
||
|
|
"error": str(e),
|
||
|
|
}
|
||
|
|
|
||
|
|
|
||
|
|
@app.get("/info")
|
||
|
|
async def get_platform_info() -> Dict[str, Any]:
|
||
|
|
"""플랫폼 정보."""
|
||
|
|
try:
|
||
|
|
adapter = await get_neo4j_adapter()
|
||
|
|
components = await get_components()
|
||
|
|
llm_manager = await get_llm_manager()
|
||
|
|
|
||
|
|
# Graph stats
|
||
|
|
analytics = components.get("analytics")
|
||
|
|
try:
|
||
|
|
stats = await analytics.get_graph_statistics() if analytics else {}
|
||
|
|
except:
|
||
|
|
stats = {}
|
||
|
|
|
||
|
|
return {
|
||
|
|
"platform": "Ontology System Construction Platform",
|
||
|
|
"phase": "7 (LLM Integration)",
|
||
|
|
"version": "0.7.0",
|
||
|
|
"components": {
|
||
|
|
"neo4j": "ok" if adapter else "unavailable",
|
||
|
|
"entity_resolver": "ok" if components.get("resolver") else "unavailable",
|
||
|
|
"subgraph_retriever": "ok" if components.get("retriever") else "unavailable",
|
||
|
|
"pattern_matcher": "ok" if components.get("matcher") else "unavailable",
|
||
|
|
"graph_analytics": "ok" if components.get("analytics") else "unavailable",
|
||
|
|
"llm_manager": "ok" if llm_manager else "unavailable",
|
||
|
|
},
|
||
|
|
"graph_stats": stats,
|
||
|
|
"llm_config": {
|
||
|
|
"provider": llm_manager.config.provider.value,
|
||
|
|
"model": llm_manager.config.model,
|
||
|
|
"temperature": llm_manager.config.temperature,
|
||
|
|
"max_tokens": llm_manager.config.max_tokens,
|
||
|
|
},
|
||
|
|
}
|
||
|
|
except Exception as e:
|
||
|
|
logger.error(f"Info retrieval failed: {e}")
|
||
|
|
raise HTTPException(status_code=500, detail=str(e))
|
||
|
|
|
||
|
|
|
||
|
|
# ============================================================================
|
||
|
|
# Router Registration
|
||
|
|
# ============================================================================
|
||
|
|
|
||
|
|
app.include_router(llm_router)
|
||
|
|
|
||
|
|
|
||
|
|
if __name__ == "__main__":
|
||
|
|
import uvicorn
|
||
|
|
uvicorn.run(app, host="0.0.0.0", port=8001)
|