File size: 3,941 Bytes
fba6023
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
from __future__ import annotations

import re
from enum import Enum
from typing import Any

from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator

IDENTIFIER_PATTERN = r"^[a-z][a-z0-9_]*$"


class ParameterType(str, Enum):
    """Supported runtime parameter types for YAML templates."""

    STRING = "string"
    INTEGER = "integer"
    NUMBER = "number"
    BOOLEAN = "boolean"
    ARRAY = "array"
    OBJECT = "object"


class ParameterDefinition(BaseModel):
    """Declarative validation contract for one runtime template parameter."""

    model_config = ConfigDict(extra="forbid")

    type: ParameterType
    description: str = ""
    required: bool = Field(default=False, strict=True)
    default: Any = None
    enum: list[Any] | None = None
    minimum: float | None = None
    maximum: float | None = None
    min_length: int | None = Field(default=None, ge=0)
    max_length: int | None = Field(default=None, ge=0)

    @property
    def has_default(self) -> bool:
        """Return whether YAML explicitly supplied a default value."""
        return "default" in self.model_fields_set

    @model_validator(mode="after")
    def validate_bounds(self) -> ParameterDefinition:
        if self.minimum is not None and self.maximum is not None:
            if self.minimum > self.maximum:
                raise ValueError("minimum must not exceed maximum")
        if self.min_length is not None and self.max_length is not None:
            if self.min_length > self.max_length:
                raise ValueError("min_length must not exceed max_length")
        return self


class PipelineStep(BaseModel):
    """One operation invocation in a template pipeline."""

    model_config = ConfigDict(extra="allow")

    operation: str = Field(pattern=IDENTIFIER_PATTERN)
    when: Any = True
    inputs: list[str] | None = None
    save_as: str | None = Field(default=None, pattern=IDENTIFIER_PATTERN)

    def operation_parameters(self) -> dict[str, Any]:
        """Return operation arguments declared as extra YAML keys."""
        return dict(self.model_extra or {})


class OutputDefinition(BaseModel):
    """Declared output contract for a template."""

    model_config = ConfigDict(extra="forbid")

    format: str
    filename: str | None = None


class TemplateDefinition(BaseModel):
    """Validated, versioned YAML media workflow definition."""

    model_config = ConfigDict(extra="forbid")

    id: str = Field(pattern=IDENTIFIER_PATTERN)
    name: str = Field(min_length=1, max_length=120)
    category: str = Field(pattern=IDENTIFIER_PATTERN)
    description: str = Field(min_length=1, max_length=1000)
    author: str = Field(min_length=1, max_length=120)
    version: int = Field(ge=1, strict=True)
    tags: list[str] = Field(min_length=1)
    estimated_runtime: str = Field(min_length=1, max_length=80)
    supported_inputs: list[str] = Field(min_length=1)
    supported_outputs: list[str] = Field(min_length=1)
    parameters: dict[str, ParameterDefinition] = Field(default_factory=dict)
    pipeline: list[PipelineStep] = Field(min_length=1)
    output: OutputDefinition
    examples: list[dict[str, Any]] = Field(default_factory=list)

    @field_validator("tags", "supported_inputs", "supported_outputs")
    @classmethod
    def normalize_string_list(cls, values: list[str]) -> list[str]:
        normalized = [value.strip().lower() for value in values if value.strip()]
        if not normalized:
            raise ValueError("list must contain at least one non-empty value")
        return list(dict.fromkeys(normalized))

    @field_validator("parameters")
    @classmethod
    def validate_parameter_names(
        cls, values: dict[str, ParameterDefinition]
    ) -> dict[str, ParameterDefinition]:
        for name in values:
            if not re.fullmatch(IDENTIFIER_PATTERN, name):
                raise ValueError(f"invalid parameter name: {name}")
        return values