Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 3 additions & 1 deletion packages/gsm8k/src/aviary/envs/gsm8k/env.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,7 @@
Messages,
TaskDataset,
Tool,
ToolRequestMessage,
ToolResponseMessage,
)
from pydantic import BaseModel, ConfigDict
Expand Down Expand Up @@ -141,7 +142,8 @@ async def reset(self) -> tuple[Messages, list[Tool]]:
return [Message(content=self.problem)], self.tools

async def step(self, action: Message) -> tuple[Messages, float, bool, bool]:
action = self.check_action_is_tool_request(action)
if not isinstance(action, ToolRequestMessage):
return self.default_no_tool_calls_response
if not action.tool_calls:
return (
[
Expand Down
4 changes: 3 additions & 1 deletion packages/hotpotqa/src/aviary/envs/hotpotqa/env.py
Original file line number Diff line number Diff line change
Expand Up @@ -33,6 +33,7 @@
Messages,
TaskDataset,
Tool,
ToolRequestMessage,
eval_answer,
)
from bs4 import BeautifulSoup
Expand Down Expand Up @@ -300,7 +301,8 @@ async def step(self, action: Message) -> tuple[Messages, float, bool, bool]:
>>> print(obs, done)
[ToolResponseMessage(name='Search', tool_call_id='tool_call_id', content='...')], False
"""
action = self.check_action_is_tool_request(action)
if not isinstance(action, ToolRequestMessage):
return self.default_no_tool_calls_response
self.state.steps += 1
if not action.tool_calls:
return (
Expand Down
6 changes: 4 additions & 2 deletions packages/labbench/src/aviary/envs/labbench/env.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,7 @@
Messages,
MultipleChoiceEvaluation,
MultipleChoiceQuestion,
ToolRequestMessage,
)
from aviary.env import ENV_REGISTRY
from ldp.utils import discounted_returns
Expand Down Expand Up @@ -141,8 +142,9 @@ async def _evaluate_answer(self) -> TEvaluation:
return evaluation # type: ignore[return-value]

async def step(self, action: Message) -> tuple[Messages, float, bool, bool]:
tool_request = self.check_action_is_tool_request(action)
messages, reward, done, truncated = await super().step(tool_request)
if not isinstance(action, ToolRequestMessage):
return self.default_no_tool_calls_response
messages, reward, done, truncated = await super().step(action)
if not done or not isinstance(self._query, MultipleChoiceQuestion):
return messages, reward, done, truncated
evaluation = await self._evaluate_answer()
Expand Down
6 changes: 4 additions & 2 deletions packages/lfrqa/src/aviary/envs/lfrqa/env.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,7 @@
Message,
Messages,
MultipleChoiceQuestion,
ToolRequestMessage,
)
from aviary.env import ENV_REGISTRY
from aviary.envs.labbench import GradablePaperQAEnvironment
Expand Down Expand Up @@ -270,10 +271,11 @@ async def _evaluate_answer(self) -> dict:
return evaluation

async def step(self, action: Message) -> tuple[Messages, float, bool, bool]:
tool_request = self.check_action_is_tool_request(action)
if not isinstance(action, ToolRequestMessage):
return self.default_no_tool_calls_response
messages, reward, done, truncated = await super(
GradablePaperQAEnvironment, self
).step(tool_request)
).step(action)
if not done:
return messages, reward, done, truncated
evaluation = await self._evaluate_answer()
Expand Down
9 changes: 4 additions & 5 deletions packages/notebook/src/aviary/envs/notebook/env.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,7 +11,7 @@

import aiodocker
import nbformat
from aviary.core import Environment, Message, Messages, Tool
from aviary.core import Environment, Message, Messages, Tool, ToolRequestMessage
from aviary.message import EnvStateMessage
from jupyter_client.manager import AsyncKernelManager
from nbformat import NotebookNode
Expand Down Expand Up @@ -184,14 +184,13 @@ async def reset(self) -> tuple[Messages, list[Tool]]:
return init_obs, self.tools

async def step(self, action: Message) -> tuple[Messages, float, bool, bool]:
tool_request = self.check_action_is_tool_request(action)
if not isinstance(action, ToolRequestMessage):
return self.default_no_tool_calls_response
prev_reward = self.state.total_reward

obs = cast(
Messages,
await self.exec_tool_calls(
tool_request, concurrency=False, handle_tool_exc=True
),
await self.exec_tool_calls(action, concurrency=False, handle_tool_exc=True),
)
reward = self.state.total_reward - prev_reward

Expand Down
24 changes: 15 additions & 9 deletions src/aviary/env.py
Original file line number Diff line number Diff line change
Expand Up @@ -103,13 +103,18 @@ class Environment(ABC, Generic[TEnvState]):
tools: list[Tool]
state: TEnvState

@staticmethod
def check_action_is_tool_request(action: Message) -> ToolRequestMessage:
if not isinstance(action, ToolRequestMessage):
raise TypeError(
f"Expected a ToolRequestMessage, but got {type(action).__name__}."
)
return action
@property
def default_no_tool_calls_response(self) -> tuple[Messages, float, bool, bool]:
return (
[
Message(
content="No tool calls received. Please call one or more tools to proceed."
)
],
0.0,
False,
False,
)

Comment thread
sidnarayanan marked this conversation as resolved.
@abstractmethod
async def step(self, action: Message) -> tuple[Messages, float, bool, bool]:
Expand Down Expand Up @@ -552,9 +557,10 @@ def from_task(cls, task: str) -> "DummyEnv":
return cls(task=task)

async def step(self, action: Message) -> tuple[Messages, float, bool, bool]:
tool_request = self.check_action_is_tool_request(action)
if not isinstance(action, ToolRequestMessage):
return self.default_no_tool_calls_response
msgs: Messages = await self.exec_tool_calls( # type: ignore[assignment]
tool_request, state=self.state, concurrency=self.concurrent_tool_calls
action, state=self.state, concurrency=self.concurrent_tool_calls
) or [
ToolResponseMessage(
content=f"No tool calls input in tool request {action}.",
Expand Down
7 changes: 4 additions & 3 deletions src/aviary/functional.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,7 +6,7 @@

from aviary.env import Environment, Frame
from aviary.message import Message
from aviary.tools import Messages, Tool
from aviary.tools import Messages, Tool, ToolRequestMessage
from aviary.utils import is_coroutine_callable


Expand Down Expand Up @@ -72,9 +72,10 @@ async def reset(self) -> tuple[Messages, list[Tool]]:
return [Message(content=obs)], self.tools

async def step(self, action: Message) -> tuple[Messages, float, bool, bool]:
tool_request = self.check_action_is_tool_request(action)
if not isinstance(action, ToolRequestMessage):
return self.default_no_tool_calls_response
msgs = await self.exec_tool_calls(
tool_request,
action,
state=self.state,
concurrency=self.allow_concurrency,
handle_tool_exc=True,
Expand Down
45 changes: 37 additions & 8 deletions tests/test_envs.py
Original file line number Diff line number Diff line change
Expand Up @@ -353,8 +353,9 @@ def slow_tool() -> None:
return [], self.tools

async def step(self, action: Message) -> tuple[Messages, float, bool, bool]:
tool_request = self.check_action_is_tool_request(action)
await self.exec_tool_calls(tool_request, exec_timeout=0.0001)
if not isinstance(action, ToolRequestMessage):
return self.default_no_tool_calls_response
await self.exec_tool_calls(action, exec_timeout=0.0001)

return [], 0.0, False, False

Expand Down Expand Up @@ -754,9 +755,12 @@ async def test_step_with_plain_message(self, server_async_client: AsyncClient):
response = await server_async_client.post(
"/step", json={"env_id": env_id, "action": action.model_dump()}
)
assert response.status_code == 500, (
"Plain Message should fail check_action_is_tool_request"
)
assert response.status_code == 200
obs, reward, done, truncated = response.json()
assert any("No tool calls" in m["content"] for m in obs)
assert reward == 0.0
assert not done
assert not truncated

@pytest.mark.asyncio
async def test_step_with_tool_response_message(
Expand All @@ -772,9 +776,34 @@ async def test_step_with_tool_response_message(
response = await server_async_client.post(
"/step", json={"env_id": env_id, "action": action.model_dump()}
)
assert response.status_code == 500, (
"ToolResponseMessage should fail check_action_is_tool_request"
)
assert response.status_code == 200
obs, reward, done, truncated = response.json()
assert any("No tool calls" in m["content"] for m in obs)
assert reward == 0.0
assert not done
assert not truncated


class TestDefaultNoToolCallsResponse:
@pytest.mark.parametrize(
"action",
[
Message(content="hello"),
ToolResponseMessage(content="result", name="tool", tool_call_id="abc"),
],
)
@pytest.mark.asyncio
async def test_non_tool_request_returns_default(
self, dummy_env: DummyEnv, action: Message
):
await dummy_env.reset()
obs, reward, done, truncated = await dummy_env.step(action)
assert len(obs) == 1
assert obs[0].content
assert "No tool calls" in obs[0].content
assert reward == 0.0
assert not done
assert not truncated


class TestStepRequestDeserialization:
Expand Down
Loading