92 lines
2.8 KiB
Python
92 lines
2.8 KiB
Python
from __future__ import annotations
|
|
|
|
from skills_learning.config import SkillAuthoringConfig
|
|
from skills_learning.embeddings import embed_skill_text
|
|
from skills_learning.models import FlowStep, FlowTemplateSkill, SkillMetadata
|
|
from skills_learning.retrieval import retrieve_candidate_skills
|
|
from skills_learning.store import SkillStore
|
|
|
|
|
|
class FakeEmbeddingClient:
|
|
def __init__(self, *, fail: bool = False) -> None:
|
|
self.fail = fail
|
|
|
|
def embed(self, text: str, *, model: str) -> list[float]:
|
|
if self.fail:
|
|
raise TimeoutError("embedding timeout")
|
|
lowered = text.lower()
|
|
if "coffee" in lowered:
|
|
return [1.0, 0.0]
|
|
if "tea" in lowered:
|
|
return [0.0, 1.0]
|
|
return [0.5, 0.5]
|
|
|
|
|
|
def _skill(name: str, goal: str) -> FlowTemplateSkill:
|
|
return FlowTemplateSkill(
|
|
metadata=SkillMetadata(
|
|
name=name,
|
|
description=f"Learned flow for {goal}",
|
|
originating_goal=goal,
|
|
),
|
|
steps=[FlowStep("input_text", {"text": goal})],
|
|
parameters={},
|
|
)
|
|
|
|
|
|
def test_embed_skill_text_degrades_to_none_on_provider_failure() -> None:
|
|
result = embed_skill_text(
|
|
"coffee",
|
|
client=FakeEmbeddingClient(fail=True),
|
|
config=SkillAuthoringConfig(enabled=True),
|
|
)
|
|
|
|
assert result is None
|
|
|
|
|
|
def test_retrieve_candidate_skills_ranks_by_similarity_and_truncates_top_k() -> None:
|
|
store = SkillStore()
|
|
coffee = store.create_version(_skill("search coffee", "search coffee"))
|
|
tea = store.create_version(_skill("search tea", "search tea"))
|
|
store.store_embedding(coffee, [1.0, 0.0], model_name="fake")
|
|
store.store_embedding(tea, [0.0, 1.0], model_name="fake")
|
|
|
|
results = retrieve_candidate_skills(
|
|
"find coffee",
|
|
store=store,
|
|
top_k=1,
|
|
embedding_client=FakeEmbeddingClient(),
|
|
config=SkillAuthoringConfig(enabled=True),
|
|
)
|
|
|
|
assert [result.skill.name for result in results] == ["search coffee"]
|
|
|
|
|
|
def test_retrieve_candidate_skills_returns_empty_without_embeddings() -> None:
|
|
store = SkillStore()
|
|
store.create_version(_skill("search coffee", "search coffee"))
|
|
|
|
results = retrieve_candidate_skills(
|
|
"find coffee",
|
|
store=store,
|
|
embedding_client=FakeEmbeddingClient(),
|
|
config=SkillAuthoringConfig(enabled=True),
|
|
)
|
|
|
|
assert results == []
|
|
|
|
|
|
def test_skill_without_embedding_is_stored_but_excluded_from_retrieval() -> None:
|
|
store = SkillStore()
|
|
skill = store.create_version(_skill("search coffee", "search coffee"))
|
|
|
|
results = retrieve_candidate_skills(
|
|
"find coffee",
|
|
store=store,
|
|
embedding_client=FakeEmbeddingClient(),
|
|
config=SkillAuthoringConfig(enabled=True),
|
|
)
|
|
|
|
assert store.get_by_id(skill.id) == skill
|
|
assert results == []
|