Files
2026-07-06 23:15:30 +08:00

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,
)