mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
Fix substring matching specificity and remove mutable Reducto OCR config state
- Fireworks: _get_model_cost_capability fallback now picks the longest substring match in model_cost so more specific entries win over less specific ones (instead of returning the first match by insertion order). - Reducto OCR: drop per-request _api_key/_api_base instance attributes on _BaseReductoOCRConfig and instead thread api_key/api_base through transform_ocr_request/async_transform_ocr_request kwargs from the shared OCR HTTP handler. Makes the config safe to share/cache across concurrent requests with different credentials. Co-authored-by: Yassin Kortam <yassin@berri.ai>
This commit is contained in:
parent
69b5a0788a
commit
f71ab0f1a1
3 changed files with 74 additions and 28 deletions
|
|
@ -1409,6 +1409,8 @@ class BaseLLMHTTPHandler:
|
|||
document=document,
|
||||
optional_params=optional_params,
|
||||
headers=headers,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
)
|
||||
|
||||
# All providers return OCRRequestData
|
||||
|
|
@ -1477,6 +1479,8 @@ class BaseLLMHTTPHandler:
|
|||
document=document,
|
||||
optional_params=optional_params,
|
||||
headers=headers,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
)
|
||||
|
||||
# All providers return OCRRequestData
|
||||
|
|
|
|||
|
|
@ -272,8 +272,11 @@ class FireworksAIConfig(OpenAIGPTConfig):
|
|||
|
||||
# Fallback: preserve historical substring matching for model name
|
||||
# variants (e.g. fine-tuned or regionally-suffixed versions of a
|
||||
# known model). Look for any fireworks_ai entry whose normalized
|
||||
# short name is a substring of our normalized short name.
|
||||
# known model). Pick the *longest* matching entry so a more specific
|
||||
# known model (e.g. "qwen3-8b-instruct") wins over a less specific
|
||||
# one (e.g. "qwen3-8b") when the query model is more specific still.
|
||||
best_match_short: Optional[str] = None
|
||||
best_match_value: Optional[bool] = None
|
||||
for key, model_info in litellm.model_cost.items():
|
||||
if not key.startswith("fireworks_ai/"):
|
||||
continue
|
||||
|
|
@ -284,10 +287,13 @@ class FireworksAIConfig(OpenAIGPTConfig):
|
|||
key_short = key[len("fireworks_ai/") :]
|
||||
if key_short.startswith("accounts/fireworks/models/"):
|
||||
key_short = key_short[len("accounts/fireworks/models/") :]
|
||||
if key_short and key_short in short_name:
|
||||
return cast(Optional[bool], model_info.get(capability))
|
||||
if not key_short or key_short not in short_name:
|
||||
continue
|
||||
if best_match_short is None or len(key_short) > len(best_match_short):
|
||||
best_match_short = key_short
|
||||
best_match_value = cast(Optional[bool], model_info.get(capability))
|
||||
|
||||
return None
|
||||
return best_match_value
|
||||
|
||||
def get_provider_info(self, model: str) -> ProviderSpecificModelInfo:
|
||||
supports_function_calling_value = self._get_model_cost_capability(
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
from typing import Any, Dict, Optional
|
||||
from typing import Any, Dict, Optional, Tuple
|
||||
|
||||
import httpx
|
||||
|
||||
|
|
@ -19,11 +19,6 @@ from litellm.llms.reducto.common import (
|
|||
|
||||
|
||||
class _BaseReductoOCRConfig(BaseOCRConfig):
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
self._api_key: Optional[str] = None
|
||||
self._api_base: Optional[str] = None
|
||||
|
||||
def map_ocr_params(
|
||||
self,
|
||||
non_default_params: dict,
|
||||
|
|
@ -54,9 +49,6 @@ class _BaseReductoOCRConfig(BaseOCRConfig):
|
|||
"Missing REDUCTO_API_KEY - set it in the environment or pass api_key to litellm.ocr()/litellm.aocr()"
|
||||
)
|
||||
|
||||
self._api_key = resolved_key
|
||||
self._api_base = (api_base or REDUCTO_API_BASE).rstrip("/")
|
||||
|
||||
return {
|
||||
"Authorization": f"Bearer {resolved_key}",
|
||||
"Content-Type": "application/json",
|
||||
|
|
@ -83,32 +75,56 @@ class _BaseReductoOCRConfig(BaseOCRConfig):
|
|||
)
|
||||
return source_url
|
||||
|
||||
def _ensure_file_id_sync(self, model: str, document: DocumentType) -> str:
|
||||
@staticmethod
|
||||
def _resolve_credentials(
|
||||
api_key: Optional[str], api_base: Optional[str]
|
||||
) -> Tuple[str, str]:
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
|
||||
resolved_key = api_key or get_secret_str("REDUCTO_API_KEY")
|
||||
if resolved_key is None:
|
||||
raise ValueError(
|
||||
"Missing REDUCTO_API_KEY - set it in the environment or pass api_key to litellm.ocr()/litellm.aocr()"
|
||||
)
|
||||
resolved_base = (api_base or REDUCTO_API_BASE).rstrip("/")
|
||||
return resolved_key, resolved_base
|
||||
|
||||
def _ensure_file_id_sync(
|
||||
self,
|
||||
model: str,
|
||||
document: DocumentType,
|
||||
api_key: Optional[str],
|
||||
api_base: Optional[str],
|
||||
) -> str:
|
||||
source_url = self._get_source_url(document=document, model=model)
|
||||
file_id, raw_bytes, mime = extract_file_id_or_bytes(source_url, model=model)
|
||||
if file_id is not None:
|
||||
return file_id
|
||||
if self._api_key is None:
|
||||
raise ValueError("Reducto API key was not initialized before OCR upload.")
|
||||
resolved_key, resolved_base = self._resolve_credentials(api_key, api_base)
|
||||
return upload_bytes_sync(
|
||||
raw_bytes=raw_bytes or b"",
|
||||
mime=mime,
|
||||
api_key=self._api_key,
|
||||
api_base=self._api_base,
|
||||
api_key=resolved_key,
|
||||
api_base=resolved_base,
|
||||
)
|
||||
|
||||
async def _ensure_file_id_async(self, model: str, document: DocumentType) -> str:
|
||||
async def _ensure_file_id_async(
|
||||
self,
|
||||
model: str,
|
||||
document: DocumentType,
|
||||
api_key: Optional[str],
|
||||
api_base: Optional[str],
|
||||
) -> str:
|
||||
source_url = self._get_source_url(document=document, model=model)
|
||||
file_id, raw_bytes, mime = extract_file_id_or_bytes(source_url, model=model)
|
||||
if file_id is not None:
|
||||
return file_id
|
||||
if self._api_key is None:
|
||||
raise ValueError("Reducto API key was not initialized before OCR upload.")
|
||||
resolved_key, resolved_base = self._resolve_credentials(api_key, api_base)
|
||||
return await upload_bytes_async(
|
||||
raw_bytes=raw_bytes or b"",
|
||||
mime=mime,
|
||||
api_key=self._api_key,
|
||||
api_base=self._api_base,
|
||||
api_key=resolved_key,
|
||||
api_base=resolved_base,
|
||||
)
|
||||
|
||||
def transform_ocr_response(
|
||||
|
|
@ -146,7 +162,12 @@ class ReductoParseV3Config(_BaseReductoOCRConfig):
|
|||
headers: dict,
|
||||
**kwargs,
|
||||
) -> OCRRequestData:
|
||||
file_id = self._ensure_file_id_sync(model=model, document=document)
|
||||
file_id = self._ensure_file_id_sync(
|
||||
model=model,
|
||||
document=document,
|
||||
api_key=kwargs.get("api_key"),
|
||||
api_base=kwargs.get("api_base"),
|
||||
)
|
||||
return OCRRequestData(data={"input": file_id, **optional_params}, files=None)
|
||||
|
||||
async def async_transform_ocr_request(
|
||||
|
|
@ -157,7 +178,12 @@ class ReductoParseV3Config(_BaseReductoOCRConfig):
|
|||
headers: dict,
|
||||
**kwargs,
|
||||
) -> OCRRequestData:
|
||||
file_id = await self._ensure_file_id_async(model=model, document=document)
|
||||
file_id = await self._ensure_file_id_async(
|
||||
model=model,
|
||||
document=document,
|
||||
api_key=kwargs.get("api_key"),
|
||||
api_base=kwargs.get("api_base"),
|
||||
)
|
||||
return OCRRequestData(data={"input": file_id, **optional_params}, files=None)
|
||||
|
||||
|
||||
|
|
@ -180,7 +206,12 @@ class ReductoParseLegacyConfig(_BaseReductoOCRConfig):
|
|||
headers: dict,
|
||||
**kwargs,
|
||||
) -> OCRRequestData:
|
||||
file_id = self._ensure_file_id_sync(model=model, document=document)
|
||||
file_id = self._ensure_file_id_sync(
|
||||
model=model,
|
||||
document=document,
|
||||
api_key=kwargs.get("api_key"),
|
||||
api_base=kwargs.get("api_base"),
|
||||
)
|
||||
return OCRRequestData(
|
||||
data=self._build_legacy_body(
|
||||
file_id=file_id, optional_params=optional_params
|
||||
|
|
@ -196,7 +227,12 @@ class ReductoParseLegacyConfig(_BaseReductoOCRConfig):
|
|||
headers: dict,
|
||||
**kwargs,
|
||||
) -> OCRRequestData:
|
||||
file_id = await self._ensure_file_id_async(model=model, document=document)
|
||||
file_id = await self._ensure_file_id_async(
|
||||
model=model,
|
||||
document=document,
|
||||
api_key=kwargs.get("api_key"),
|
||||
api_base=kwargs.get("api_base"),
|
||||
)
|
||||
return OCRRequestData(
|
||||
data=self._build_legacy_body(
|
||||
file_id=file_id, optional_params=optional_params
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue