from __future__ import annotations
from abc import ABC, ABCMeta, abstractmethod
from typing import TYPE_CHECKING, Any, Union
if TYPE_CHECKING:
from areal.api.engine_api import InferenceEngine
from areal.experimental.openai.types import InteractionWithTokenLogpReward
class RolloutWorkflow(ABC):
@abstractmethod
async def arun_episode(
self, engine: InferenceEngine, data: dict[str, Any]
) -> dict[str, Any] | None | dict[str, InteractionWithTokenLogpReward]:
"""Run a single episode of the workflow.
Note
----
Returning `None` implies that this trajectory is rejected and will not be used for training.
See concrete example implementations under the `areal/workflow` directory.
Parameters
----------
engine : InferenceEngine
The inference engine to use for generating responses
data : Dict[str, Any]
Input data for the workflow episode
Returns
-------
Dict[str, Any] | None | Dict[str, InteractionWithTokenLogpReward]
The trajectory result, None if rejected, or a dictionary of completion results
"""
raise NotImplementedError()
class _DeprecatedAgentWorkflowMeta(ABCMeta):
"""Metaclass that ensures deprecation warning triggers on any subclass instantiation.
This approach guarantees the warning fires even if subclasses forget to call
super().__init__(), since __call__ executes before any __init__ method.
Inherits from ABCMeta to maintain compatibility with ABC.
"""
def __call__(cls, *args, **kwargs):
import warnings
warnings.warn(
f"{cls.__name__} inherits from deprecated AgentWorkflow. "
"You no longer need to inherit from this class. "
"Any class with a compatible async run() method will work.",
DeprecationWarning,
stacklevel=2,
)
return super().__call__(*args, **kwargs)
class AgentWorkflow(ABC, metaclass=_DeprecatedAgentWorkflowMeta):
"""Base class for agent-based workflows (DEPRECATED).
.. deprecated:: 1.0.0
Inheriting from AgentWorkflow is no longer required. Any class with
a compatible ``run()`` method will work. This class is kept for
backward compatibility but may be removed in a future version.
To use agent-based workflows, simply implement a class with::
async def run(self, data: dict[str, Any], **extra_kwargs: Any) -> dict[str, float] | float
"""
@abstractmethod
async def run(
self, data: dict[str, Any], **extra_kwargs: Any
) -> dict[str, float] | float:
"""Run an agent with any SDK, e.g., OpenAI SDK.
`data` contains the input data for the agent, which may
include any parameters or hyperparameters required.
`extra_kwargs` includes parameters provided by AReaL:
- base_url: str
The base URL of the OpenAI-compatible proxy server
- http_client: httpx.AsyncClient
The HTTP client to use for making requests in AsyncOpenAI
- api_key: str
The session-scoped API key for authenticating with the proxy
Parameters
----------
data : dict[str, Any]
Input data for the agent workflow
Returns
-------
dict[str, float] | float
The final reward or a dictionary of reward keyed by response ID
"""
raise NotImplementedError()
WorkflowLike = Union[
"RolloutWorkflow",
type["RolloutWorkflow"],
str,
Any,
]