workflow
This commit is contained in:
@@ -0,0 +1,128 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
from typing import Protocol, Sequence
|
||||
|
||||
from core.models import Scene
|
||||
from workflow.models import ConditionSpec, WorkflowStepResult
|
||||
|
||||
|
||||
class UnknownConditionKindError(Exception):
|
||||
pass
|
||||
|
||||
|
||||
class ConditionEvaluator(Protocol):
|
||||
def evaluate(
|
||||
self,
|
||||
spec: ConditionSpec,
|
||||
*,
|
||||
scene: Scene | None,
|
||||
world_state: object | None,
|
||||
started_at: datetime | None = None,
|
||||
step_results: Sequence[WorkflowStepResult] = (),
|
||||
) -> bool:
|
||||
...
|
||||
|
||||
|
||||
class SceneContainsTextEvaluator:
|
||||
def evaluate(
|
||||
self,
|
||||
spec: ConditionSpec,
|
||||
*,
|
||||
scene: Scene | None,
|
||||
world_state: object | None,
|
||||
started_at: datetime | None = None,
|
||||
step_results: Sequence[WorkflowStepResult] = (),
|
||||
) -> bool:
|
||||
if scene is None:
|
||||
return False
|
||||
target = str(spec.params.get("text") or spec.params.get("target") or "")
|
||||
if not target:
|
||||
return False
|
||||
target_lower = target.lower()
|
||||
return any(
|
||||
target_lower in str(element.text or "").lower()
|
||||
for element in scene.elements
|
||||
)
|
||||
|
||||
|
||||
class WorldVariableEqualsEvaluator:
|
||||
def evaluate(
|
||||
self,
|
||||
spec: ConditionSpec,
|
||||
*,
|
||||
scene: Scene | None,
|
||||
world_state: object | None,
|
||||
started_at: datetime | None = None,
|
||||
step_results: Sequence[WorkflowStepResult] = (),
|
||||
) -> bool:
|
||||
variables = getattr(world_state, "variables", None)
|
||||
if not isinstance(variables, dict):
|
||||
return False
|
||||
name = spec.params.get("name")
|
||||
if not isinstance(name, str):
|
||||
return False
|
||||
return variables.get(name) == spec.params.get("value")
|
||||
|
||||
|
||||
class ElapsedSecondsEvaluator:
|
||||
def evaluate(
|
||||
self,
|
||||
spec: ConditionSpec,
|
||||
*,
|
||||
scene: Scene | None,
|
||||
world_state: object | None,
|
||||
started_at: datetime | None = None,
|
||||
step_results: Sequence[WorkflowStepResult] = (),
|
||||
) -> bool:
|
||||
if started_at is None:
|
||||
return False
|
||||
seconds = float(spec.params.get("seconds") or 0)
|
||||
return (datetime.now(started_at.tzinfo) - started_at).total_seconds() >= seconds
|
||||
|
||||
|
||||
class StepResultSuccessEvaluator:
|
||||
def evaluate(
|
||||
self,
|
||||
spec: ConditionSpec,
|
||||
*,
|
||||
scene: Scene | None,
|
||||
world_state: object | None,
|
||||
started_at: datetime | None = None,
|
||||
step_results: Sequence[WorkflowStepResult] = (),
|
||||
) -> bool:
|
||||
target_step_id = spec.params.get("step_id")
|
||||
return any(
|
||||
result.step_id == target_step_id and result.success
|
||||
for result in step_results
|
||||
)
|
||||
|
||||
|
||||
DEFAULT_CONDITION_REGISTRY: dict[str, ConditionEvaluator] = {
|
||||
"scene_contains_text": SceneContainsTextEvaluator(),
|
||||
"world_variable_equals": WorldVariableEqualsEvaluator(),
|
||||
"elapsed_seconds": ElapsedSecondsEvaluator(),
|
||||
"step_result_success": StepResultSuccessEvaluator(),
|
||||
}
|
||||
|
||||
|
||||
def evaluate_condition(
|
||||
spec: ConditionSpec,
|
||||
*,
|
||||
scene: Scene | None,
|
||||
world_state: object | None,
|
||||
started_at: datetime | None = None,
|
||||
step_results: Sequence[WorkflowStepResult] = (),
|
||||
registry: dict[str, ConditionEvaluator] | None = None,
|
||||
) -> bool:
|
||||
evaluators = registry or DEFAULT_CONDITION_REGISTRY
|
||||
evaluator = evaluators.get(spec.kind)
|
||||
if evaluator is None:
|
||||
raise UnknownConditionKindError(spec.kind)
|
||||
return evaluator.evaluate(
|
||||
spec,
|
||||
scene=scene,
|
||||
world_state=world_state,
|
||||
started_at=started_at,
|
||||
step_results=step_results,
|
||||
)
|
||||
Reference in New Issue
Block a user