Files

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