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 )