Files
ai-agent-book/chapter2/prompt-engineering/tau_bench/envs/base.py
T
liqiang b119135836
Build latest book artifacts / build (push) Canceled after 0s
dependency resolution / resolve (3.11) (push) Canceled after 0s
dependency resolution / resolve (3.13) (push) Canceled after 0s
deploy-pages / build (push) Canceled after 0s
deploy-pages / deploy (push) Canceled after 0s
i18n consistency check / check (push) Canceled after 0s
provider adoption tests / test (chapter2/context-compression) (push) Canceled after 0s
provider adoption tests / test (chapter2/prompt-injection) (push) Canceled after 0s
provider adoption tests / test (chapter2/system-hint) (push) Canceled after 0s
provider adoption tests / test (chapter3/log-sanitization) (push) Canceled after 0s
web-search-agent tests / test (push) Canceled after 0s
web-search-agent tests / agentbook (push) Canceled after 0s
ai-agent-book 精选快照(<2MB 代码与文档,来自 github.com/bojieli/ai-agent-book)
2026-08-20 13:12:50 +00:00

167 lines
5.6 KiB
Python

# Copyright Sierra
import random
from hashlib import sha256
from tau_bench.envs.tool import Tool
from typing import Any, Callable, Dict, List, Type, Optional, Set, Union, Tuple
from tau_bench.envs.user import load_user, UserStrategy
from tau_bench.types import (
Action,
Task,
EnvInfo,
EnvResetResponse,
EnvResponse,
RewardResult,
RewardOutputInfo,
RewardActionInfo,
RESPOND_ACTION_NAME,
)
ToHashable = Union[
str, int, float, Dict[str, "ToHashable"], List["ToHashable"], Set["ToHashable"]
]
Hashable = Union[str, int, float, Tuple["Hashable"], Tuple[Tuple[str, "Hashable"]]]
def to_hashable(item: ToHashable) -> Hashable:
if isinstance(item, dict):
return tuple((key, to_hashable(value)) for key, value in sorted(item.items()))
elif isinstance(item, list):
return tuple(to_hashable(element) for element in item)
elif isinstance(item, set):
return tuple(sorted(to_hashable(element) for element in item))
else:
return item
def consistent_hash(
value: Hashable,
) -> str:
return sha256(str(value).encode("utf-8")).hexdigest()
class Env(object):
def __init__(
self,
data_load_func: Callable[[], Dict[str, Any]],
tools: List[Type[Tool]],
tasks: List[Task],
wiki: str,
rules: List[str],
user_strategy: Union[str, UserStrategy],
user_model: str,
user_provider: Optional[str] = None,
task_index: Optional[int] = None,
user_seed: Optional[int] = None,
) -> None:
super().__init__()
self.data_load_func = data_load_func
self.data = data_load_func()
self.tools_map: Dict[str, Type[Tool]] = {
tool.get_info()["function"]["name"]: tool for tool in tools
}
self.tools_info = [tool.get_info() for tool in tools]
self.terminate_tools = []
self.tasks = tasks
if task_index is not None:
self.task_index = task_index
else:
self.task_index = random.randrange(len(tasks))
self.task = tasks[self.task_index]
self.wiki = wiki
self.rules = rules
self.user = load_user(
user_strategy=user_strategy, model=user_model, provider=user_provider,
seed=user_seed,
)
self.actions: List[Action] = []
def reset(self, task_index: Optional[int] = None) -> EnvResetResponse:
if task_index is None:
task_index = random.randrange(len(self.tasks))
self.task_index = task_index
self.data = self.data_load_func()
self.task = self.tasks[task_index]
self.actions = []
initial_observation = self.user.reset(instruction=self.task.instruction)
return EnvResetResponse(
observation=initial_observation, info=EnvInfo(task=self.task, source="user")
)
def step(self, action: Action) -> EnvResponse:
self.actions.append(action)
info = EnvInfo(task=self.task)
reward = 0
done = False
if action.name == RESPOND_ACTION_NAME:
observation = self.user.step(action.kwargs["content"])
info.source = "user"
done = "###STOP###" in observation
elif action.name in self.tools_map:
try:
observation = self.tools_map[action.name].invoke(
data=self.data, **action.kwargs
)
except Exception as e:
observation = f"Error: {e}"
info.source = action.name
if action.name in self.terminate_tools:
done = True
else:
observation = f"Unknown action {action.name}"
info.source = action.name
if done:
reward_res = self.calculate_reward()
reward = reward_res.reward
info.reward_info = reward_res
info.user_cost = self.user.get_total_cost()
return EnvResponse(observation=observation, reward=reward, done=done, info=info)
def get_data_hash(self) -> str:
return consistent_hash(to_hashable(self.data))
def calculate_reward(self) -> RewardResult:
data_hash = self.get_data_hash()
reward = 1.0
actions = [
action for action in self.task.actions if action.name != RESPOND_ACTION_NAME
]
# Check if the database changes are correct. If they are not correct, then we set the reward to 0.
# TODO: cache gt_data_hash in tasks.py (low priority)
self.data = self.data_load_func()
for action in self.task.actions:
if action.name not in self.terminate_tools:
self.step(action)
gt_data_hash = self.get_data_hash()
info = RewardActionInfo(
r_actions=data_hash == gt_data_hash, gt_data_hash=gt_data_hash
)
if not info.r_actions:
reward = 0.0
if len(self.task.outputs) > 0:
# check outputs
r_outputs = 1.0
outputs = {}
for output in self.task.outputs:
found = False
for action in self.actions:
if (
action.name == RESPOND_ACTION_NAME
and output.lower()
in action.kwargs["content"].lower().replace(",", "")
):
found = True
break
outputs[output] = found
if not found:
r_outputs = 0.0
reward = 0.0
info = RewardOutputInfo(r_outputs=r_outputs, outputs=outputs)
return RewardResult(reward=reward, info=info, actions=actions)