ai-agent-book 精选快照(<2MB 代码与文档,来自 github.com/bojieli/ai-agent-book)
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

This commit is contained in:
2026-08-20 13:12:50 +00:00
commit b119135836
10275 changed files with 3284984 additions and 0 deletions
@@ -0,0 +1,432 @@
from __future__ import annotations
import argparse
from typing import Any, TypeVar
from pydantic import BaseModel
from tau_bench.model_utils.api._model_methods import MODEL_METHODS
from tau_bench.model_utils.api.cache import cache_call_w_dedup
from tau_bench.model_utils.api.datapoint import (
BinaryClassifyDatapoint,
ClassifyDatapoint,
Datapoint,
GenerateDatapoint,
ParseDatapoint,
ParseForceDatapoint,
ScoreDatapoint,
)
from tau_bench.model_utils.api.logging import log_call
from tau_bench.model_utils.api.router import RequestRouter, default_request_router
from tau_bench.model_utils.api.sample import (
EnsembleSamplingStrategy,
MajoritySamplingStrategy,
SamplingStrategy,
get_default_sampling_strategy,
)
from tau_bench.model_utils.api.types import PartialObj
from tau_bench.model_utils.model.general_model import GeneralModel
from tau_bench.model_utils.model.model import (
AnyModel,
BinaryClassifyModel,
ClassifyModel,
GenerateModel,
ParseForceModel,
ParseModel,
ScoreModel,
)
T = TypeVar("T", bound=BaseModel)
class API(object):
wrappers_for_main_methods = [log_call, cache_call_w_dedup]
def __init__(
self,
parse_models: list[ParseModel],
generate_models: list[GenerateModel],
parse_force_models: list[ParseForceModel],
score_models: list[ScoreModel],
classify_models: list[ClassifyModel],
binary_classify_models: list[BinaryClassifyModel] | None = None,
sampling_strategy: SamplingStrategy | None = None,
request_router: RequestRouter | None = None,
log_file: str | None = None,
) -> None:
if sampling_strategy is None:
sampling_strategy = get_default_sampling_strategy()
if request_router is None:
request_router = default_request_router()
self.sampling_strategy = sampling_strategy
self.request_router = request_router
self._log_file = log_file
self.binary_classify_models = binary_classify_models
self.classify_models = classify_models
self.parse_models = parse_models
self.generate_models = generate_models
self.parse_force_models = parse_force_models
self.score_models = score_models
self.__init_subclass__()
self.__init_subclass__()
def __init_subclass__(cls):
for method_name in MODEL_METHODS:
if hasattr(cls, method_name):
method = getattr(cls, method_name)
for wrapper in cls.wrappers_for_main_methods:
method = wrapper(method)
setattr(cls, method_name, method)
@classmethod
def from_general_model(
cls,
model: GeneralModel,
sampling_strategy: SamplingStrategy | None = None,
request_router: RequestRouter | None = None,
log_file: str | None = None,
) -> "API":
return cls(
binary_classify_models=[model],
classify_models=[model],
parse_models=[model],
generate_models=[model],
parse_force_models=[model],
score_models=[model],
log_file=log_file,
sampling_strategy=sampling_strategy,
request_router=request_router,
)
@classmethod
def from_general_models(
cls,
models: list[GeneralModel],
sampling_strategy: SamplingStrategy | None = None,
request_router: RequestRouter | None = None,
log_file: str | None = None,
) -> "API":
if len(models) == 0:
raise ValueError("Must provide at least one model")
return cls(
binary_classify_models=models,
classify_models=models,
parse_models=models,
generate_models=models,
parse_force_models=models,
score_models=models,
log_file=log_file,
sampling_strategy=sampling_strategy,
request_router=request_router,
)
def set_default_binary_classify_models(self, models: list[BinaryClassifyModel]) -> None:
if len(models) == 0:
raise ValueError("Must provide at least one model")
self.binary_classify_models = models
def set_default_classify_models(self, models: list[BinaryClassifyModel]) -> None:
if len(models) == 0:
raise ValueError("Must provide at least one model")
self.classify_models = models
def set_default_parse_models(self, models: list[ParseModel]) -> None:
if len(models) == 0:
raise ValueError("Must provide at least one model")
self.parse_models = models
def set_default_generate_models(self, models: list[GenerateModel]) -> None:
if len(models) == 0:
raise ValueError("Must provide at least one model")
self.generate_models = models
def set_default_parse_force_models(self, models: list[ParseForceModel]) -> None:
if len(models) == 0:
raise ValueError("Must provide at least one model")
self.parse_force_models = models
def set_default_score_models(self, models: list[ScoreModel]) -> None:
if len(models) == 0:
raise ValueError("Must provide at least one model")
self.score_models = models
def set_default_sampling_strategy(self, sampling_strategy: SamplingStrategy) -> None:
self.sampling_strategy = sampling_strategy
def set_default_request_router(self, request_router: RequestRouter) -> None:
self.request_router = request_router
def _run_with_sampling_strategy(
self,
models: list[AnyModel],
datapoint: Datapoint,
sampling_strategy: SamplingStrategy,
) -> T:
assert len(models) > 0
def _run_datapoint(model: AnyModel, temp: float | None = None) -> T:
if isinstance(datapoint, ClassifyDatapoint):
return model.classify(
instruction=datapoint.instruction,
text=datapoint.text,
options=datapoint.options,
examples=datapoint.examples,
temperature=temp,
)
elif isinstance(datapoint, BinaryClassifyDatapoint):
return model.binary_classify(
instruction=datapoint.instruction,
text=datapoint.text,
examples=datapoint.examples,
temperature=temp,
)
elif isinstance(datapoint, ParseForceDatapoint):
return model.parse_force(
instruction=datapoint.instruction,
typ=datapoint.typ,
text=datapoint.text,
examples=datapoint.examples,
temperature=temp,
)
elif isinstance(datapoint, GenerateDatapoint):
return model.generate(
instruction=datapoint.instruction,
text=datapoint.text,
examples=datapoint.examples,
temperature=temp,
)
elif isinstance(datapoint, ParseDatapoint):
return model.parse(
text=datapoint.text,
typ=datapoint.typ,
examples=datapoint.examples,
temperature=temp,
)
elif isinstance(datapoint, ScoreDatapoint):
return model.score(
instruction=datapoint.instruction,
text=datapoint.text,
min=datapoint.min,
max=datapoint.max,
examples=datapoint.examples,
temperature=temp,
)
else:
raise ValueError(f"Unknown datapoint type: {type(datapoint)}")
if isinstance(sampling_strategy, EnsembleSamplingStrategy):
return sampling_strategy.execute(
[lambda x=model: _run_datapoint(x, 0.0) for model in models]
)
return sampling_strategy.execute(
lambda: _run_datapoint(
models[0], 0.2 if isinstance(sampling_strategy, MajoritySamplingStrategy) else None
)
)
def _api_call(
self, models: list[AnyModel], datapoint: Datapoint, sampling_strategy: SamplingStrategy
) -> T:
if isinstance(sampling_strategy, EnsembleSamplingStrategy):
return self._run_with_sampling_strategy(models, datapoint, sampling_strategy)
model = self.request_router.route(dp=datapoint, available_models=models)
return self._run_with_sampling_strategy(
models=[model], datapoint=datapoint, sampling_strategy=sampling_strategy
)
def classify(
self,
instruction: str,
text: str,
options: list[str],
examples: list[ClassifyDatapoint] | None = None,
sampling_strategy: SamplingStrategy | None = None,
request_router: RequestRouter | None = None,
models: list[ClassifyModel] | None = None,
) -> int:
if models is None:
models = self.classify_models
if sampling_strategy is None:
sampling_strategy = self.sampling_strategy
if request_router is None:
request_router = self.request_router
return self._api_call(
models=models,
datapoint=ClassifyDatapoint(
instruction=instruction, text=text, options=options, examples=examples
),
sampling_strategy=sampling_strategy,
)
def binary_classify(
self,
instruction: str,
text: str,
examples: list[BinaryClassifyDatapoint] | None = None,
sampling_strategy: SamplingStrategy | None = None,
request_router: RequestRouter | None = None,
models: list[BinaryClassifyModel] | None = None,
) -> bool:
if models is None:
models = (
self.binary_classify_models
if self.binary_classify_models is not None
else self.classify_models
)
if sampling_strategy is None:
sampling_strategy = self.sampling_strategy
if request_router is None:
request_router = self.request_router
return self._api_call(
models=models,
datapoint=BinaryClassifyDatapoint(
instruction=instruction, text=text, examples=examples
),
sampling_strategy=sampling_strategy,
)
def parse(
self,
text: str,
typ: type[T] | dict[str, Any],
examples: list[ParseDatapoint] | None = None,
sampling_strategy: SamplingStrategy | None = None,
request_router: RequestRouter | None = None,
models: list[ParseModel] | None = None,
) -> T | PartialObj | dict[str, Any]:
if models is None:
models = self.parse_models
if sampling_strategy is None:
sampling_strategy = self.sampling_strategy
if request_router is None:
request_router = self.request_router
return self._api_call(
models=models,
datapoint=ParseDatapoint(text=text, typ=typ, examples=examples),
sampling_strategy=sampling_strategy,
)
def generate(
self,
instruction: str,
text: str,
examples: list[GenerateDatapoint] | None = None,
sampling_strategy: SamplingStrategy | None = None,
request_router: RequestRouter | None = None,
models: list[GenerateModel] | None = None,
) -> str:
if models is None:
models = self.generate_models
if sampling_strategy is None:
sampling_strategy = self.sampling_strategy
if request_router is None:
request_router = self.request_router
return self._api_call(
models=models,
datapoint=GenerateDatapoint(instruction=instruction, text=text, examples=examples),
sampling_strategy=sampling_strategy,
)
def parse_force(
self,
instruction: str,
typ: type[T] | dict[str, Any],
text: str | None = None,
examples: list[ParseForceDatapoint] | None = None,
sampling_strategy: SamplingStrategy | None = None,
request_router: RequestRouter | None = None,
models: list[ParseForceModel] | None = None,
) -> T | dict[str, Any]:
if models is None:
models = self.parse_force_models
if sampling_strategy is None:
sampling_strategy = self.sampling_strategy
if request_router is None:
request_router = self.request_router
return self._api_call(
models=models,
datapoint=ParseForceDatapoint(
instruction=instruction, typ=typ, text=text, examples=examples
),
sampling_strategy=sampling_strategy,
)
def score(
self,
instruction: str,
text: str,
min: int,
max: int,
examples: list[ScoreDatapoint] | None = None,
sampling_strategy: SamplingStrategy | None = None,
request_router: RequestRouter | None = None,
models: list[ScoreModel] | None = None,
) -> int:
if models is None:
models = self.score_models
if sampling_strategy is None:
sampling_strategy = self.sampling_strategy
if request_router is None:
request_router = self.request_router
return self._api_call(
models=models,
datapoint=ScoreDatapoint(
instruction=instruction, text=text, min=min, max=max, examples=examples
),
sampling_strategy=sampling_strategy,
)
def default_api(
log_file: str | None = None,
sampling_strategy: SamplingStrategy | None = None,
request_router: RequestRouter | None = None,
) -> API:
from tau_bench.model_utils.model.general_model import default_model
model = default_model()
return API(
binary_classify_models=[model],
classify_models=[model],
parse_models=[model],
generate_models=[model],
parse_force_models=[model],
score_models=[model],
sampling_strategy=sampling_strategy,
request_router=request_router,
log_file=log_file,
)
def default_api_from_args(args: argparse.Namespace) -> API:
from tau_bench.model_utils.model.general_model import model_factory
model = model_factory(model_id=args.model, platform=args.platform, base_url=args.base_url)
return API.from_general_model(model=model)
def default_quick_api(
log_file: str | None = None,
sampling_strategy: SamplingStrategy | None = None,
request_router: RequestRouter | None = None,
) -> API:
from tau_bench.model_utils.model.general_model import default_quick_model
model = default_quick_model()
return API(
binary_classify_models=[model],
classify_models=[model],
parse_models=[model],
generate_models=[model],
parse_force_models=[model],
score_models=[model],
sampling_strategy=sampling_strategy,
request_router=request_router,
log_file=log_file,
)