mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
Merge db52939f75 into b781d157d7
This commit is contained in:
commit
8394975fff
6 changed files with 127 additions and 4 deletions
1
.github/workflows/cost-map-guard.yml
vendored
1
.github/workflows/cost-map-guard.yml
vendored
|
|
@ -39,5 +39,6 @@ jobs:
|
|||
MERGE_BASE: ${{ steps.revisions.outputs.merge_base }}
|
||||
HEAD_SHA: ${{ github.event.pull_request.head.sha }}
|
||||
HEAD_REF: ${{ github.event.pull_request.head.ref }}
|
||||
LITELLM_LOCAL_MODEL_COST_MAP: "True"
|
||||
run: |
|
||||
uv run --frozen python ci_cd/cost_map_guard.py --base "$MERGE_BASE" --head "$HEAD_SHA" --head-ref "$HEAD_REF"
|
||||
|
|
|
|||
|
|
@ -1,8 +1,10 @@
|
|||
"""Guard the cost map on pull requests.
|
||||
|
||||
Every pull request whose diff against its merge base touches one of the three cost map files gets the file
|
||||
checks: the files parse, the backup copy matches the root file, and the JSON schema is in sync and validates the
|
||||
map. A pull request that leaves all three untouched skips them, since merging it keeps the base branch's copies
|
||||
checks: the files parse, the backup copy matches the root file, the JSON schema is in sync and validates the
|
||||
map, and every Bedrock regional or cross-region row mirrors the supports_* keys of its base row (the runtime
|
||||
resolves capabilities per row, and tests/local_testing/test_get_model_info.py fails the base branch on any
|
||||
drift). A pull request that leaves all three untouched skips them, since merging it keeps the base branch's copies
|
||||
and its head tree only carries whatever state the branch was cut from. Pull requests from the cost map sync bot
|
||||
(branches named litellm_cost_map_sync_*) always get the file checks and additionally may only touch those three
|
||||
files and may only add or update models.
|
||||
|
|
@ -12,21 +14,27 @@ from __future__ import annotations
|
|||
|
||||
import argparse
|
||||
import json
|
||||
import re
|
||||
import subprocess
|
||||
import sys
|
||||
from collections.abc import Sequence
|
||||
from collections.abc import Iterator, Mapping, Sequence
|
||||
from dataclasses import dataclass
|
||||
from typing import Final
|
||||
|
||||
from generate_model_prices_schema import SPECIAL_ROOT_KEYS, build_schema, render, validation_errors
|
||||
|
||||
from litellm.llms.bedrock.common_utils import get_bedrock_base_model
|
||||
|
||||
COST_MAP_PATH: Final = "model_prices_and_context_window.json"
|
||||
BACKUP_PATH: Final = "litellm/model_prices_and_context_window_backup.json"
|
||||
SCHEMA_PATH: Final = "model_prices_and_context_window.schema.json"
|
||||
GUARDED_PATHS: Final = (COST_MAP_PATH, BACKUP_PATH, SCHEMA_PATH)
|
||||
BOT_BRANCH_PREFIX: Final = "litellm_cost_map_sync_"
|
||||
BEDROCK_COMMITMENT_PREFIX: Final = re.compile(r"[136]-month-commitment/")
|
||||
BEDROCK_PARITY_RULE: Final = "a Bedrock regional or cross-region row mirrors every supports_* key of its base row"
|
||||
|
||||
CostMap = dict[str, object]
|
||||
Entries = Mapping[str, Mapping[str, object]]
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
|
|
@ -75,6 +83,7 @@ def _file_failures(head: Snapshot, head_map: CostMap) -> tuple[str, ...]:
|
|||
f"{COST_MAP_PATH} does not validate against its schema: {error}"
|
||||
for error in validation_errors(head_map, json.loads(schema_text))[:20]
|
||||
),
|
||||
*_bedrock_parity_failures(_entries(head_map)),
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -82,6 +91,46 @@ def _entries(cost_map: CostMap) -> dict[str, dict[str, object]]:
|
|||
return {key: entry for key, entry in cost_map.items() if isinstance(entry, dict)}
|
||||
|
||||
|
||||
def _is_bedrock_row(key: str, entry: Mapping[str, object]) -> bool:
|
||||
return str(entry.get("litellm_provider", "")).startswith("bedrock") and "invoke/" not in key
|
||||
|
||||
|
||||
def _bedrock_base_key(key: str, entries: Entries) -> str | None:
|
||||
candidate: Final = BEDROCK_COMMITMENT_PREFIX.sub("", key.replace("*/", ""))
|
||||
base: Final = get_bedrock_base_model(candidate)
|
||||
if base == candidate:
|
||||
return None
|
||||
base_key: Final = base if base in entries else f"bedrock/{base}"
|
||||
return base_key if base_key in entries and base_key != key else None
|
||||
|
||||
|
||||
def _capability_gaps(
|
||||
key: str, entry: Mapping[str, object], base_key: str, base_entry: Mapping[str, object]
|
||||
) -> Iterator[str]:
|
||||
for capability, base_value in base_entry.items():
|
||||
if not capability.startswith("supports_"):
|
||||
continue
|
||||
if capability not in entry:
|
||||
yield (
|
||||
f"{key} lacks {capability} ({json.dumps(base_value)}) carried by its base row {base_key}; "
|
||||
f"{BEDROCK_PARITY_RULE}"
|
||||
)
|
||||
elif entry[capability] != base_value:
|
||||
yield (
|
||||
f"{key}.{capability} is {json.dumps(entry[capability])} but {base_key}.{capability} is "
|
||||
f"{json.dumps(base_value)}; {BEDROCK_PARITY_RULE}"
|
||||
)
|
||||
|
||||
|
||||
def _bedrock_parity_failures(entries: Entries) -> Iterator[str]:
|
||||
for key, entry in entries.items():
|
||||
if not _is_bedrock_row(key, entry):
|
||||
continue
|
||||
base_key = _bedrock_base_key(key, entries)
|
||||
if base_key is not None:
|
||||
yield from _capability_gaps(key, entry, base_key, entries[base_key])
|
||||
|
||||
|
||||
def _bot_failures(base: Snapshot, head_map: CostMap, changed_files: Sequence[str]) -> tuple[str, ...]:
|
||||
base_map: Final = _parse_object(base.cost_map, COST_MAP_PATH)
|
||||
if isinstance(base_map, str):
|
||||
|
|
|
|||
|
|
@ -2492,6 +2492,7 @@
|
|||
"supports_native_structured_output": false,
|
||||
"supports_max_reasoning_effort": true,
|
||||
"supports_minimal_reasoning_effort": true,
|
||||
"supports_output_config": true,
|
||||
"supports_parallel_tool_use_config": true,
|
||||
"prompt_cache_min_tokens": 2048,
|
||||
"source": "https://aws.amazon.com/bedrock/pricing/"
|
||||
|
|
|
|||
|
|
@ -2492,6 +2492,7 @@
|
|||
"supports_native_structured_output": false,
|
||||
"supports_max_reasoning_effort": true,
|
||||
"supports_minimal_reasoning_effort": true,
|
||||
"supports_output_config": true,
|
||||
"supports_parallel_tool_use_config": true,
|
||||
"prompt_cache_min_tokens": 2048,
|
||||
"source": "https://aws.amazon.com/bedrock/pricing/"
|
||||
|
|
|
|||
|
|
@ -323,10 +323,12 @@ def test_get_model_info_bedrock_cross_region_capability_parity():
|
|||
Cross-region inference profiles carry litellm_provider "bedrock_converse", so the
|
||||
regional drift check above (which filters on "bedrock") never reaches them.
|
||||
"""
|
||||
from litellm.llms.bedrock.common_utils import get_bedrock_cross_region_inference_regions
|
||||
|
||||
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
|
||||
litellm.model_cost = litellm.get_model_cost_map(url="")
|
||||
|
||||
prefixes = ("us.", "eu.", "apac.", "us-gov.")
|
||||
prefixes = tuple(f"{region}." for region in get_bedrock_cross_region_inference_regions())
|
||||
checked = 0
|
||||
|
||||
for k, v in litellm.model_cost.items():
|
||||
|
|
|
|||
|
|
@ -175,6 +175,75 @@ def test_bot_may_not_change_special_root_keys() -> None:
|
|||
assert _failures(head) == ("bot PRs may not change fallback_generalizations",)
|
||||
|
||||
|
||||
def _bedrock_entry(**capabilities: bool) -> dict[str, object]:
|
||||
return _entry(litellm_provider="bedrock_converse", **capabilities)
|
||||
|
||||
|
||||
BEDROCK_REGIONAL_ROWS: Final = (
|
||||
("us-gov.nvidia.nemotron-nano-3-30b", "nvidia.nemotron-nano-3-30b"),
|
||||
("jp.anthropic.claude-opus-4-7", "anthropic.claude-opus-4-7"),
|
||||
("bedrock/ap-northeast-1/deepseek.v3.2", "bedrock/deepseek.v3.2"),
|
||||
("bedrock/*/1-month-commitment/cohere.command-text-v14", "bedrock/cohere.command-text-v14"),
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(("regional_key", "base_key"), BEDROCK_REGIONAL_ROWS)
|
||||
def test_bedrock_regional_row_missing_a_base_capability_is_reported(regional_key: str, base_key: str) -> None:
|
||||
head = _snapshot(
|
||||
{
|
||||
**BASE_MAP,
|
||||
base_key: _bedrock_entry(supports_audio_input=False, supports_vision=False),
|
||||
regional_key: _bedrock_entry(supports_vision=False),
|
||||
}
|
||||
)
|
||||
assert _failures(head, bot=False) == (
|
||||
f"{regional_key} lacks supports_audio_input (false) carried by its base row {base_key}; "
|
||||
f"{guard.BEDROCK_PARITY_RULE}",
|
||||
)
|
||||
|
||||
|
||||
def test_bedrock_cross_region_row_contradicting_its_base_is_reported() -> None:
|
||||
head = _snapshot(
|
||||
{
|
||||
**BASE_MAP,
|
||||
"anthropic.claude-mythos-preview": _bedrock_entry(supports_prompt_caching=False),
|
||||
"us.anthropic.claude-mythos-preview": _bedrock_entry(supports_prompt_caching=True),
|
||||
}
|
||||
)
|
||||
assert _failures(head, bot=False) == (
|
||||
"us.anthropic.claude-mythos-preview.supports_prompt_caching is true but "
|
||||
f"anthropic.claude-mythos-preview.supports_prompt_caching is false; {guard.BEDROCK_PARITY_RULE}",
|
||||
)
|
||||
|
||||
|
||||
def test_bedrock_parity_is_checked_for_bots_too() -> None:
|
||||
head = _snapshot(
|
||||
{
|
||||
**BASE_MAP,
|
||||
"deepseek.v3.2": _bedrock_entry(supports_vision=False),
|
||||
"us.deepseek.v3.2": _bedrock_entry(),
|
||||
}
|
||||
)
|
||||
assert _failures(head, bot=True) == (
|
||||
f"us.deepseek.v3.2 lacks supports_vision (false) carried by its base row deepseek.v3.2; "
|
||||
f"{guard.BEDROCK_PARITY_RULE}",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"rows",
|
||||
[
|
||||
pytest.param({"us.deepseek.v3.2": _bedrock_entry(supports_vision=False, supports_pdf_input=True)}, id="extra"),
|
||||
pytest.param({"bedrock/invoke/deepseek.v3.2": _bedrock_entry()}, id="invoke"),
|
||||
pytest.param({"eu.orphan.model": _bedrock_entry()}, id="no-base"),
|
||||
pytest.param({"openrouter/us.deepseek.v3.2": _entry()}, id="other-provider"),
|
||||
],
|
||||
)
|
||||
def test_bedrock_rows_in_parity_or_outside_the_rule_pass(rows: dict[str, dict[str, object]]) -> None:
|
||||
head = _snapshot({**BASE_MAP, "deepseek.v3.2": _bedrock_entry(supports_vision=False), **rows})
|
||||
assert _failures(head, bot=False) == ()
|
||||
|
||||
|
||||
def _commit(repo: Path, cost_map: dict[str, object], message: str) -> str:
|
||||
text = _serialize(cost_map)
|
||||
(repo / guard.COST_MAP_PATH).write_text(text)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue