workflow
This commit is contained in:
@@ -0,0 +1,244 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import sqlite3
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from core.models import utc_now
|
||||
from workflow.config import DEFAULT_WORKFLOW_DB_PATH
|
||||
from workflow.models import (
|
||||
WorkflowDefinition,
|
||||
WorkflowRun,
|
||||
WorkflowRunStatus,
|
||||
WorkflowStepResult,
|
||||
)
|
||||
|
||||
|
||||
class WorkflowStore:
|
||||
def __init__(self, db_path: str | Path = DEFAULT_WORKFLOW_DB_PATH) -> None:
|
||||
self.db_path = Path(db_path)
|
||||
self.db_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
self._ensure_schema()
|
||||
|
||||
def save_definition(self, definition: WorkflowDefinition) -> None:
|
||||
with self._connect() as connection:
|
||||
connection.execute(
|
||||
"""
|
||||
insert into workflow_definitions (id, name, definition_json)
|
||||
values (?, ?, ?)
|
||||
on conflict(id) do update set
|
||||
name = excluded.name,
|
||||
definition_json = excluded.definition_json
|
||||
""",
|
||||
(
|
||||
definition.id,
|
||||
definition.name,
|
||||
json.dumps(definition.to_dict(), ensure_ascii=False),
|
||||
),
|
||||
)
|
||||
|
||||
def get_definition(self, definition_id: str) -> WorkflowDefinition | None:
|
||||
with self._connect() as connection:
|
||||
row = connection.execute(
|
||||
"select definition_json from workflow_definitions where id = ?",
|
||||
(definition_id,),
|
||||
).fetchone()
|
||||
if row is None:
|
||||
return None
|
||||
return WorkflowDefinition.from_dict(json.loads(row["definition_json"]))
|
||||
|
||||
def create_run(
|
||||
self,
|
||||
definition_id: str,
|
||||
initial_variables: dict[str, Any] | None = None,
|
||||
*,
|
||||
device_id: str | None = None,
|
||||
) -> WorkflowRun:
|
||||
definition = self.get_definition(definition_id)
|
||||
if definition is None:
|
||||
raise KeyError(f"unknown workflow definition {definition_id}")
|
||||
run = WorkflowRun(
|
||||
definition_id=definition_id,
|
||||
status="running",
|
||||
current_step_id=definition.entry_step_id,
|
||||
variables=dict(initial_variables or {}),
|
||||
device_id=device_id,
|
||||
)
|
||||
with self._connect() as connection:
|
||||
connection.execute(
|
||||
"""
|
||||
insert into workflow_runs (
|
||||
id, definition_id, status, current_step_id, device_id,
|
||||
variables_json, created_at, updated_at
|
||||
) values (?, ?, ?, ?, ?, ?, ?, ?)
|
||||
""",
|
||||
(
|
||||
run.id,
|
||||
run.definition_id,
|
||||
run.status,
|
||||
run.current_step_id,
|
||||
run.device_id,
|
||||
json.dumps(run.variables, ensure_ascii=False),
|
||||
run.created_at.isoformat(),
|
||||
run.updated_at.isoformat(),
|
||||
),
|
||||
)
|
||||
return run
|
||||
|
||||
def append_step_result(
|
||||
self,
|
||||
run_id: str,
|
||||
step_result: WorkflowStepResult,
|
||||
) -> None:
|
||||
with self._connect() as connection:
|
||||
index = (
|
||||
connection.execute(
|
||||
"select count(*) as count from workflow_step_results where run_id = ?",
|
||||
(run_id,),
|
||||
).fetchone()["count"]
|
||||
+ 1
|
||||
)
|
||||
connection.execute(
|
||||
"""
|
||||
insert into workflow_step_results (
|
||||
run_id, step_index, step_id, kind, success, detail_json,
|
||||
task_id, timestamp
|
||||
) values (?, ?, ?, ?, ?, ?, ?, ?)
|
||||
""",
|
||||
(
|
||||
run_id,
|
||||
index,
|
||||
step_result.step_id,
|
||||
step_result.kind,
|
||||
1 if step_result.success else 0,
|
||||
json.dumps(step_result.detail, ensure_ascii=False),
|
||||
step_result.task_id,
|
||||
step_result.timestamp.isoformat(),
|
||||
),
|
||||
)
|
||||
|
||||
def update_run(
|
||||
self,
|
||||
run_id: str,
|
||||
*,
|
||||
status: WorkflowRunStatus | None = None,
|
||||
current_step_id: str | None = None,
|
||||
variables: dict[str, Any] | None = None,
|
||||
) -> None:
|
||||
run = self.get_run(run_id)
|
||||
if run is None:
|
||||
raise KeyError(f"unknown workflow run {run_id}")
|
||||
next_status = status or run.status
|
||||
next_step_id = current_step_id
|
||||
if current_step_id is None and next_status not in {
|
||||
"completed",
|
||||
"failed",
|
||||
"cancelled",
|
||||
}:
|
||||
next_step_id = run.current_step_id
|
||||
next_variables = dict(run.variables if variables is None else variables)
|
||||
with self._connect() as connection:
|
||||
connection.execute(
|
||||
"""
|
||||
update workflow_runs
|
||||
set status = ?,
|
||||
current_step_id = ?,
|
||||
variables_json = ?,
|
||||
updated_at = ?
|
||||
where id = ?
|
||||
""",
|
||||
(
|
||||
next_status,
|
||||
next_step_id,
|
||||
json.dumps(next_variables, ensure_ascii=False),
|
||||
utc_now().isoformat(),
|
||||
run_id,
|
||||
),
|
||||
)
|
||||
|
||||
def get_run(self, run_id: str) -> WorkflowRun | None:
|
||||
with self._connect() as connection:
|
||||
run_row = connection.execute(
|
||||
"select * from workflow_runs where id = ?",
|
||||
(run_id,),
|
||||
).fetchone()
|
||||
if run_row is None:
|
||||
return None
|
||||
result_rows = connection.execute(
|
||||
"""
|
||||
select * from workflow_step_results
|
||||
where run_id = ?
|
||||
order by step_index
|
||||
""",
|
||||
(run_id,),
|
||||
).fetchall()
|
||||
return WorkflowRun.from_dict(
|
||||
{
|
||||
"id": run_row["id"],
|
||||
"definition_id": run_row["definition_id"],
|
||||
"status": run_row["status"],
|
||||
"current_step_id": run_row["current_step_id"],
|
||||
"device_id": run_row["device_id"],
|
||||
"variables": json.loads(run_row["variables_json"]),
|
||||
"created_at": run_row["created_at"],
|
||||
"updated_at": run_row["updated_at"],
|
||||
"step_results": [
|
||||
{
|
||||
"step_id": row["step_id"],
|
||||
"kind": row["kind"],
|
||||
"success": bool(row["success"]),
|
||||
"detail": json.loads(row["detail_json"]),
|
||||
"task_id": row["task_id"],
|
||||
"timestamp": row["timestamp"],
|
||||
}
|
||||
for row in result_rows
|
||||
],
|
||||
}
|
||||
)
|
||||
|
||||
def _ensure_schema(self) -> None:
|
||||
with self._connect() as connection:
|
||||
connection.execute(
|
||||
"""
|
||||
create table if not exists workflow_definitions (
|
||||
id text primary key,
|
||||
name text not null,
|
||||
definition_json text not null
|
||||
)
|
||||
"""
|
||||
)
|
||||
connection.execute(
|
||||
"""
|
||||
create table if not exists workflow_runs (
|
||||
id text primary key,
|
||||
definition_id text not null,
|
||||
status text not null,
|
||||
current_step_id text,
|
||||
device_id text,
|
||||
variables_json text not null,
|
||||
created_at text not null,
|
||||
updated_at text not null
|
||||
)
|
||||
"""
|
||||
)
|
||||
connection.execute(
|
||||
"""
|
||||
create table if not exists workflow_step_results (
|
||||
id integer primary key autoincrement,
|
||||
run_id text not null,
|
||||
step_index integer not null,
|
||||
step_id text not null,
|
||||
kind text not null,
|
||||
success integer not null,
|
||||
detail_json text not null,
|
||||
task_id text,
|
||||
timestamp text not null
|
||||
)
|
||||
"""
|
||||
)
|
||||
|
||||
def _connect(self) -> sqlite3.Connection:
|
||||
connection = sqlite3.connect(self.db_path)
|
||||
connection.row_factory = sqlite3.Row
|
||||
return connection
|
||||
Reference in New Issue
Block a user