142 lines
4.3 KiB
Python
142 lines
4.3 KiB
Python
from __future__ import annotations
|
|
|
|
from dataclasses import dataclass, field, replace
|
|
from datetime import datetime
|
|
from typing import Any
|
|
from uuid import uuid4
|
|
|
|
from core.models import utc_now
|
|
from skills_learning.models import (
|
|
LOCAL_SYNTHESIS_SOURCE,
|
|
FlowTemplateSkill,
|
|
clone_skill,
|
|
)
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class SkillEmbeddingRecord:
|
|
skill_id: str
|
|
version: int
|
|
vector: list[float]
|
|
model_name: str
|
|
updated_at: datetime = field(default_factory=utc_now)
|
|
|
|
def to_dict(self) -> dict[str, Any]:
|
|
return {
|
|
"skill_id": self.skill_id,
|
|
"version": self.version,
|
|
"vector": list(self.vector),
|
|
"model_name": self.model_name,
|
|
"updated_at": self.updated_at.isoformat(),
|
|
}
|
|
|
|
|
|
class SkillStore:
|
|
def __init__(self) -> None:
|
|
self._skills: dict[str, FlowTemplateSkill] = {}
|
|
self._embeddings: dict[tuple[str, int], SkillEmbeddingRecord] = {}
|
|
|
|
def create_version(
|
|
self,
|
|
skill: FlowTemplateSkill,
|
|
*,
|
|
parent: FlowTemplateSkill | None = None,
|
|
) -> FlowTemplateSkill:
|
|
version = parent.version + 1 if parent else self._next_version(skill.name)
|
|
stored = skill.with_metadata(
|
|
id=uuid4().hex,
|
|
source=LOCAL_SYNTHESIS_SOURCE,
|
|
version=version,
|
|
parent_version_id=parent.id if parent else skill.parent_version_id,
|
|
created_at=utc_now(),
|
|
updated_at=utc_now(),
|
|
)
|
|
self._skills[stored.id] = clone_skill(stored)
|
|
return clone_skill(stored)
|
|
|
|
def update_skill(self, skill: FlowTemplateSkill) -> FlowTemplateSkill:
|
|
if skill.id not in self._skills:
|
|
raise KeyError(f"unknown skill {skill.id}")
|
|
stored = skill.with_metadata(
|
|
source=LOCAL_SYNTHESIS_SOURCE,
|
|
updated_at=utc_now(),
|
|
)
|
|
self._skills[stored.id] = clone_skill(stored)
|
|
return clone_skill(stored)
|
|
|
|
def get_by_id(self, skill_id: str) -> FlowTemplateSkill | None:
|
|
skill = self._skills.get(skill_id)
|
|
return clone_skill(skill) if skill else None
|
|
|
|
def get_latest_by_name(self, name: str) -> FlowTemplateSkill | None:
|
|
versions = self.list_versions(name)
|
|
return versions[-1] if versions else None
|
|
|
|
def list_versions(self, name: str) -> list[FlowTemplateSkill]:
|
|
return sorted(
|
|
[
|
|
clone_skill(skill)
|
|
for skill in self._skills.values()
|
|
if skill.name == name
|
|
],
|
|
key=lambda skill: skill.version,
|
|
)
|
|
|
|
def list_latest(self) -> list[FlowTemplateSkill]:
|
|
latest: dict[str, FlowTemplateSkill] = {}
|
|
for skill in self._skills.values():
|
|
current = latest.get(skill.name)
|
|
if current is None or skill.version > current.version:
|
|
latest[skill.name] = skill
|
|
return [clone_skill(skill) for skill in latest.values()]
|
|
|
|
def list_all(self) -> list[FlowTemplateSkill]:
|
|
return [clone_skill(skill) for skill in self._skills.values()]
|
|
|
|
def store_embedding(
|
|
self,
|
|
skill: FlowTemplateSkill,
|
|
vector: list[float],
|
|
*,
|
|
model_name: str,
|
|
) -> SkillEmbeddingRecord:
|
|
record = SkillEmbeddingRecord(
|
|
skill_id=skill.id,
|
|
version=skill.version,
|
|
vector=[float(value) for value in vector],
|
|
model_name=model_name,
|
|
)
|
|
self._embeddings[(record.skill_id, record.version)] = record
|
|
return replace(record, vector=list(record.vector))
|
|
|
|
def get_embedding(
|
|
self,
|
|
skill_id: str,
|
|
version: int,
|
|
) -> SkillEmbeddingRecord | None:
|
|
record = self._embeddings.get((skill_id, version))
|
|
return replace(record, vector=list(record.vector)) if record else None
|
|
|
|
def list_embeddings(self) -> list[SkillEmbeddingRecord]:
|
|
return [
|
|
replace(record, vector=list(record.vector))
|
|
for record in self._embeddings.values()
|
|
]
|
|
|
|
def _next_version(self, name: str) -> int:
|
|
latest = self.get_latest_by_name(name)
|
|
return latest.version + 1 if latest else 1
|
|
|
|
|
|
_DEFAULT_STORE = SkillStore()
|
|
|
|
|
|
def get_default_store() -> SkillStore:
|
|
return _DEFAULT_STORE
|
|
|
|
|
|
def reset_default_store() -> SkillStore:
|
|
global _DEFAULT_STORE
|
|
_DEFAULT_STORE = SkillStore()
|
|
return _DEFAULT_STORE
|