File size: 2,001 Bytes
5d108aa
421f883
5d108aa
421f883
 
5d108aa
 
e597175
 
 
5d108aa
421f883
 
a3dd1cb
421f883
 
 
 
 
 
5d108aa
 
421f883
5d108aa
421f883
0176e87
421f883
5d108aa
 
ba9b704
5d108aa
0176e87
 
e597175
421f883
 
0176e87
5d108aa
421f883
5d108aa
 
 
e597175
 
5d108aa
 
421f883
5d108aa
 
 
 
e597175
5d108aa
 
421f883
0176e87
5d108aa
 
421f883
0176e87
421f883
 
 
 
 
0176e87
 
 
421f883
 
 
 
 
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
"""
Server environment wrapper.

This is a thin wrapper - all logic lives in searcharena.engine.
Follows OpsArena pattern: server/ only contains the OpenEnv interface.
"""

from __future__ import annotations

from typing import Any

from openenv.core.env_server.interfaces import Environment, EnvironmentMetadata
from openenv.core.env_server.types import State

from searcharena.engine import (
    SearchEnvironment as _SearchEnvironment,
    create_sample_corpus,
    create_sample_tasks,
)
from searcharena.models import SearchAction, SearchEnvConfig, SearchObservation


class SearchEnvironment(Environment[SearchAction, SearchObservation, State]):
    """
    OpenEnv-compatible wrapper for SearchArena environment.

    This thin wrapper delegates all logic to searcharena.engine.SearchEnvironment.
    """

    SUPPORTS_CONCURRENT_SESSIONS: bool = False

    def __init__(
        self,
        config: SearchEnvConfig | None = None,
        corpus: Any | None = None,
        tasks: list | None = None,
    ):
        super().__init__()
        self._env = _SearchEnvironment(config=config, corpus=corpus, tasks=tasks)

    def reset(
        self,
        seed: int | None = None,
        episode_id: str | None = None,
        **kwargs: Any,
    ) -> SearchObservation:
        return self._env.reset(seed=seed, episode_id=episode_id, **kwargs)

    def step(
        self,
        action: SearchAction,
        timeout_s: float | None = None,
        **kwargs: Any,
    ) -> SearchObservation:
        return self._env.step(action, timeout_s=timeout_s, **kwargs)

    @property
    def state(self) -> State:
        return self._env.state

    def get_metadata(self) -> EnvironmentMetadata:
        return EnvironmentMetadata(
            name="SearchArena",
            description="Multi-hop document retrieval environment for training search agents.",
            version="0.1.0",
        )


__all__ = [
    "SearchEnvironment",
    "create_sample_corpus",
    "create_sample_tasks",
]