mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix: restore Fireworks substring matching and use RLock for Vertex sync refresh
- Fireworks _get_model_cost_capability: after exact-key lookups, fall back to substring matching against fireworks_ai/* entries in model_cost so model name variants (e.g. fine-tuned suffixes) continue to inherit capability flags like supports_reasoning. - Vertex vertex_llm_base: replace non-reentrant threading.Lock with RLock on the sync refresh path so the reauthentication retry, which recurses into get_access_token while still holding the lock, does not deadlock when reloaded credentials are also expired. Co-authored-by: Yassin Kortam <yassin@berri.ai>
This commit is contained in:
parent
acb37dc664
commit
b8050a19ba
2 changed files with 21 additions and 1 deletions
|
|
@ -270,6 +270,23 @@ class FireworksAIConfig(OpenAIGPTConfig):
|
|||
if model_info is not None and model_info.get(capability) is not None:
|
||||
return cast(Optional[bool], model_info.get(capability))
|
||||
|
||||
# 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.
|
||||
for key, model_info in litellm.model_cost.items():
|
||||
if not key.startswith("fireworks_ai/"):
|
||||
continue
|
||||
if not isinstance(model_info, dict):
|
||||
continue
|
||||
if model_info.get(capability) is None:
|
||||
continue
|
||||
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))
|
||||
|
||||
return None
|
||||
|
||||
def get_provider_info(self, model: str) -> ProviderSpecificModelInfo:
|
||||
|
|
|
|||
|
|
@ -58,7 +58,10 @@ class VertexBase:
|
|||
# Tracks in-flight background refresh tasks to avoid duplicate refreshes.
|
||||
self._background_refresh_tasks: Dict[tuple, asyncio.Task] = {}
|
||||
# Protects the sync get_access_token refresh path.
|
||||
self._sync_refresh_lock = threading.Lock()
|
||||
# Use RLock so that the reauthentication retry path (which calls
|
||||
# back into get_access_token while still holding the lock) can
|
||||
# re-acquire it without deadlocking the current thread.
|
||||
self._sync_refresh_lock = threading.RLock()
|
||||
|
||||
def get_vertex_region(self, vertex_region: Optional[str], model: str) -> str:
|
||||
import litellm
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue