Source code for mmirage.core.process.processors.batch_api.config
"""Configuration for the batch API processor in MMIRAGE."""
import logging
from dataclasses import dataclass, field
from typing import Any, Dict, Literal, Optional, Sequence
from jinja2 import Environment, meta
from mmirage.config.batch_provider import BatchProviderConfig
from mmirage.core.process.base import BaseProcessorConfig, ProcessorRegistry
from mmirage.core.process.batch.provider_resolution import (
resolve_single_provider_config,
)
from mmirage.core.process.variables import BaseVar, OutputVar
logger = logging.getLogger(__name__)
BATCH_API_PROCESSOR_TYPE = "batch_api"
env = Environment()
[docs]
@dataclass
class BatchApiProcessorConfig(BaseProcessorConfig):
"""Configuration for the batch API processor.
Provider settings are declared inline in YAML and resolved to the matching
provider config class::
processors:
- type: batch_api
provider: openai
model: gpt-4o-mini
Attributes:
provider_config: Resolved provider-specific batch configuration.
export_prompts_dir: Value of --export-prompts.
"""
type: Literal["batch_api"] = "batch_api"
provider_config: Optional[BatchProviderConfig] = None
export_prompts_dir: Optional[str] = None
[docs]
@classmethod
def from_raw(cls, data: Dict[str, Any]) -> "BatchApiProcessorConfig":
"""Build the config from a raw YAML block, dispatching on ``provider``."""
block = {key: value for key, value in data.items() if key != "type"}
return cls(
type=data["type"], provider_config=resolve_single_provider_config(block)
)
[docs]
@dataclass
class BatchApiOutputVar(OutputVar):
"""Output variable generated by the batch API processor.
Attributes:
prompt: Jinja2 template for the request prompt.
output_schema: Expected JSON field names, sent to the provider as a
hint. Unlike `LLMOutputVar`, typed/bounded field specs are not
supported here: providers are not given schema constraints, only
field names.
output_type: Output format - "JSON" or "plain".
"""
prompt: str = ""
output_schema: list[str] = field(default_factory=list)
output_type: str = ""
[docs]
def is_computable(self, vars: Sequence[BaseVar]) -> bool:
"""Check if all variables referenced in the prompt are available."""
parsed_content = env.parse(self.prompt)
template_vars = meta.find_undeclared_variables(parsed_content)
var_names = set(map(lambda v: v.name, vars))
undeclared_vars = template_vars - var_names
if len(undeclared_vars) > 0:
logger.warning(
f"Undeclared variables found for {self.name}: {undeclared_vars}"
)
return False
return True
ProcessorRegistry.register_types(
BATCH_API_PROCESSOR_TYPE, BatchApiProcessorConfig, BatchApiOutputVar
)