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:
Cursor Agent 2026-05-20 19:06:43 +00:00
parent acb37dc664
commit b8050a19ba
No known key found for this signature in database
2 changed files with 21 additions and 1 deletions

View file

@ -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:

View file

@ -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