45 lines
1.2 KiB
Python
45 lines
1.2 KiB
Python
from __future__ import annotations
|
|
|
|
from typing import Any, Protocol
|
|
|
|
from skills_learning.config import SkillAuthoringConfig, load_config
|
|
|
|
|
|
class EmbeddingClient(Protocol):
|
|
def embed(self, text: str, *, model: str) -> list[float]:
|
|
...
|
|
|
|
|
|
class OpenAIEmbeddingClient:
|
|
def __init__(self, *, transport: Any | None = None) -> None:
|
|
self._transport = transport
|
|
|
|
def embed(self, text: str, *, model: str) -> list[float]:
|
|
client = self._client()
|
|
response = client.embeddings.create(model=model, input=text)
|
|
return [float(value) for value in response.data[0].embedding]
|
|
|
|
def _client(self) -> Any:
|
|
if self._transport is not None:
|
|
return self._transport
|
|
from openai import OpenAI
|
|
|
|
self._transport = OpenAI()
|
|
return self._transport
|
|
|
|
|
|
def embed_skill_text(
|
|
text: str,
|
|
*,
|
|
client: EmbeddingClient | None = None,
|
|
config: SkillAuthoringConfig | None = None,
|
|
) -> list[float] | None:
|
|
settings = config or load_config()
|
|
if not settings.enabled:
|
|
return None
|
|
try:
|
|
embedding_client = client or OpenAIEmbeddingClient()
|
|
return embedding_client.embed(text, model=settings.embedding_model)
|
|
except Exception:
|
|
return None
|