feat: add skill learning runtime
This commit is contained in:
@@ -0,0 +1,217 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from core.models import Bounds, Scene, SceneElement, Task
|
||||
from runtime.executor import Executor, ExecutorConfig
|
||||
from runtime.planner import PlannedStep, Planner
|
||||
from runtime.task import TaskRunner, TaskRunnerConfig
|
||||
from skills_learning.config import SkillAuthoringConfig
|
||||
from skills_learning.retrieval import retrieve_candidate_skills
|
||||
from skills_learning.store import SkillStore
|
||||
from storage.artifact_store import ArtifactStore
|
||||
from storage.timeline import Timeline
|
||||
from tests.fakes import PNG_10X20
|
||||
|
||||
|
||||
class ScriptedPlanner(Planner):
|
||||
def __init__(self, steps: list[PlannedStep]) -> None:
|
||||
self.steps = steps
|
||||
|
||||
def plan(self, *, goal, scene, context):
|
||||
if len(context.step_results) >= len(self.steps):
|
||||
return []
|
||||
return [self.steps[len(context.step_results)]]
|
||||
|
||||
def goal_reached(self, *, goal, scene, context):
|
||||
return len(context.step_results) >= len(self.steps) and all(
|
||||
result.success for result in context.step_results
|
||||
)
|
||||
|
||||
|
||||
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 _scene() -> Scene:
|
||||
return Scene(
|
||||
width=10,
|
||||
height=20,
|
||||
elements=[
|
||||
SceneElement(
|
||||
id="input",
|
||||
type="input",
|
||||
text="Search",
|
||||
bounds=Bounds(1, 2, 3, 4),
|
||||
)
|
||||
],
|
||||
)
|
||||
|
||||
|
||||
def _runner(
|
||||
*,
|
||||
tmp_path,
|
||||
planner: Planner,
|
||||
timeline: Timeline | None = None,
|
||||
store: SkillStore | None = None,
|
||||
skill_config: SkillAuthoringConfig | None = None,
|
||||
embedding_client: FakeEmbeddingClient | None = None,
|
||||
on_task_succeeded=None,
|
||||
) -> TaskRunner:
|
||||
return TaskRunner(
|
||||
planner=planner,
|
||||
executor=Executor(
|
||||
tools={"input_text": lambda **kwargs: {"ok": True, **kwargs}},
|
||||
config=ExecutorConfig(max_retries=1, backoff_seconds=0),
|
||||
),
|
||||
timeline=timeline or Timeline(ArtifactStore(tmp_path / "history")),
|
||||
config=TaskRunnerConfig(max_steps=5),
|
||||
observer=lambda device_id: _scene(),
|
||||
screenshot_provider=lambda device_id: PNG_10X20,
|
||||
skill_store=store,
|
||||
skill_authoring_config=skill_config,
|
||||
skill_embedding_client=embedding_client,
|
||||
on_task_succeeded=on_task_succeeded,
|
||||
)
|
||||
|
||||
|
||||
def test_task_runner_calls_explicit_success_hook_once(tmp_path) -> None:
|
||||
calls: list[tuple[str, str, Timeline]] = []
|
||||
timeline = Timeline(ArtifactStore(tmp_path / "history"))
|
||||
runner = _runner(
|
||||
tmp_path=tmp_path,
|
||||
timeline=timeline,
|
||||
planner=ScriptedPlanner(
|
||||
[PlannedStep("input_text", "type", {"text": "coffee"})]
|
||||
),
|
||||
on_task_succeeded=lambda task_id, goal, timeline: calls.append(
|
||||
(task_id, goal, timeline)
|
||||
),
|
||||
)
|
||||
task = Task(goal="search coffee", device_id="phone")
|
||||
|
||||
result = runner.run(task)
|
||||
|
||||
assert result.status == "completed"
|
||||
assert calls == [(task.id, "search coffee", timeline)]
|
||||
|
||||
|
||||
def test_task_runner_skill_authoring_disabled_by_default_writes_no_skill(tmp_path) -> None:
|
||||
store = SkillStore()
|
||||
runner = _runner(
|
||||
tmp_path=tmp_path,
|
||||
planner=ScriptedPlanner(
|
||||
[PlannedStep("input_text", "type", {"text": "coffee"})]
|
||||
),
|
||||
store=store,
|
||||
)
|
||||
|
||||
result = runner.run(Task(goal="search coffee", device_id="phone"))
|
||||
|
||||
assert result.status == "completed"
|
||||
assert store.list_all() == []
|
||||
|
||||
|
||||
def test_task_runner_skill_authoring_enabled_stores_skill_and_embedding(tmp_path) -> None:
|
||||
store = SkillStore()
|
||||
runner = _runner(
|
||||
tmp_path=tmp_path,
|
||||
planner=ScriptedPlanner(
|
||||
[PlannedStep("input_text", "type", {"text": "coffee"})]
|
||||
),
|
||||
store=store,
|
||||
skill_config=SkillAuthoringConfig(enabled=True, embedding_model="fake"),
|
||||
embedding_client=FakeEmbeddingClient(),
|
||||
)
|
||||
|
||||
result = runner.run(Task(goal="search coffee", device_id="phone"))
|
||||
|
||||
assert result.status == "completed"
|
||||
skill = store.get_latest_by_name("search coffee")
|
||||
assert skill is not None
|
||||
assert skill.steps[0].args == {"text": "coffee"}
|
||||
assert store.get_embedding(skill.id, skill.version) is not None
|
||||
|
||||
|
||||
def test_task_runner_embedding_failure_still_stores_skill_without_embedding(tmp_path) -> None:
|
||||
store = SkillStore()
|
||||
runner = _runner(
|
||||
tmp_path=tmp_path,
|
||||
planner=ScriptedPlanner(
|
||||
[PlannedStep("input_text", "type", {"text": "coffee"})]
|
||||
),
|
||||
store=store,
|
||||
skill_config=SkillAuthoringConfig(enabled=True, embedding_model="fake"),
|
||||
embedding_client=FakeEmbeddingClient(fail=True),
|
||||
)
|
||||
|
||||
result = runner.run(Task(goal="search coffee", device_id="phone"))
|
||||
|
||||
assert result.status == "completed"
|
||||
skill = store.get_latest_by_name("search coffee")
|
||||
assert skill is not None
|
||||
assert store.get_embedding(skill.id, skill.version) is None
|
||||
|
||||
|
||||
def test_two_successful_tasks_promote_parameter_and_store_embedding(tmp_path) -> None:
|
||||
store = SkillStore()
|
||||
config = SkillAuthoringConfig(enabled=True, embedding_model="fake")
|
||||
|
||||
_runner(
|
||||
tmp_path=tmp_path,
|
||||
planner=ScriptedPlanner(
|
||||
[PlannedStep("input_text", "type coffee", {"text": "coffee"})]
|
||||
),
|
||||
store=store,
|
||||
skill_config=config,
|
||||
embedding_client=FakeEmbeddingClient(),
|
||||
).run(Task(goal="search coffee", device_id="phone"))
|
||||
|
||||
_runner(
|
||||
tmp_path=tmp_path,
|
||||
planner=ScriptedPlanner(
|
||||
[PlannedStep("input_text", "type tea", {"text": "tea"})]
|
||||
),
|
||||
store=store,
|
||||
skill_config=config,
|
||||
embedding_client=FakeEmbeddingClient(),
|
||||
).run(Task(goal="search tea", device_id="phone"))
|
||||
|
||||
skill = store.get_latest_by_name("search coffee")
|
||||
assert skill is not None
|
||||
assert skill.version == 1
|
||||
assert skill.steps[0].args == {"text": "{param_1}"}
|
||||
assert "param_1" in skill.parameters
|
||||
assert store.get_embedding(skill.id, skill.version) is not None
|
||||
|
||||
|
||||
def test_retrieve_candidate_skills_after_successful_task(tmp_path) -> None:
|
||||
store = SkillStore()
|
||||
config = SkillAuthoringConfig(enabled=True, embedding_model="fake")
|
||||
_runner(
|
||||
tmp_path=tmp_path,
|
||||
planner=ScriptedPlanner(
|
||||
[PlannedStep("input_text", "type coffee", {"text": "coffee"})]
|
||||
),
|
||||
store=store,
|
||||
skill_config=config,
|
||||
embedding_client=FakeEmbeddingClient(),
|
||||
).run(Task(goal="search coffee", device_id="phone"))
|
||||
|
||||
results = retrieve_candidate_skills(
|
||||
"find coffee",
|
||||
store=store,
|
||||
embedding_client=FakeEmbeddingClient(),
|
||||
config=config,
|
||||
)
|
||||
|
||||
assert [result.skill.name for result in results] == ["search coffee"]
|
||||
Reference in New Issue
Block a user