41 lines
1.2 KiB
Python
41 lines
1.2 KiB
Python
"""structured output 공용 래퍼.
|
|
|
|
OpenAI 호환 /chat/completions를 쓰는 프로바이더면 base_url만 바꿔 그대로 쓴다.
|
|
"""
|
|
from typing import TypeVar
|
|
|
|
from openai import AsyncOpenAI
|
|
from pydantic import BaseModel
|
|
|
|
T = TypeVar("T", bound=BaseModel)
|
|
|
|
|
|
class StructuredLLM:
|
|
def __init__(self, model: str, api_key: str, base_url: str | None = None):
|
|
self.model = model
|
|
self._client = AsyncOpenAI(api_key=api_key, base_url=base_url)
|
|
|
|
async def ask(
|
|
self,
|
|
schema: type[T],
|
|
prompt: str,
|
|
*,
|
|
system: str | None = None,
|
|
temperature: float = 0.0,
|
|
) -> T:
|
|
messages = [{"role": "system", "content": system}] if system else []
|
|
messages.append({"role": "user", "content": prompt})
|
|
|
|
msg = (await self._client.chat.completions.parse(
|
|
model=self.model,
|
|
messages=messages,
|
|
response_format=schema,
|
|
temperature=temperature,
|
|
)).choices[0].message
|
|
|
|
if msg.refusal:
|
|
raise RuntimeError(f"{self.model} 거절: {msg.refusal}")
|
|
if msg.parsed is None:
|
|
raise RuntimeError(f"{self.model} 파싱 실패: {msg.content!r}")
|
|
return msg.parsed
|