129 lines
3.6 KiB
Python
129 lines
3.6 KiB
Python
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,
|
|
)
|