fix: reduce TTFT by caching model lookups in chat completion (#20886)
fix: reduce TTFT by caching model lookups in chat completion Skip expensive get_all_models() calls when models are already cached in app.state. This significantly reduces Time To First Token (TTFT) for chat completions and embeddings requests. Previously, every request called get_all_models() which fetches model lists from all configured backends. Now we check the cache first and only call get_all_models() on cache miss. Affected endpoints: - openai: generate_chat_completion, embeddings - ollama: embed, embeddings Fixes #20069 Co-authored-by: Michael <42099345+mickeytheseal@users.noreply.github.com>
This commit is contained in:
@@ -1027,14 +1027,17 @@ async def embed(
|
||||
log.info(f"generate_ollama_batch_embeddings {form_data}")
|
||||
|
||||
if url_idx is None:
|
||||
await get_all_models(request, user=user)
|
||||
models = request.app.state.OLLAMA_MODELS
|
||||
|
||||
model = form_data.model
|
||||
|
||||
if ":" not in model:
|
||||
model = f"{model}:latest"
|
||||
|
||||
# Check if model is already in app state cache to avoid expensive get_all_models() call
|
||||
models = request.app.state.OLLAMA_MODELS
|
||||
if not models or model not in models:
|
||||
await get_all_models(request, user=user)
|
||||
models = request.app.state.OLLAMA_MODELS
|
||||
|
||||
if model in models:
|
||||
url_idx = random.choice(models[model]["urls"])
|
||||
else:
|
||||
@@ -1112,14 +1115,17 @@ async def embeddings(
|
||||
log.info(f"generate_ollama_embeddings {form_data}")
|
||||
|
||||
if url_idx is None:
|
||||
await get_all_models(request, user=user)
|
||||
models = request.app.state.OLLAMA_MODELS
|
||||
|
||||
model = form_data.model
|
||||
|
||||
if ":" not in model:
|
||||
model = f"{model}:latest"
|
||||
|
||||
# Check if model is already in app state cache to avoid expensive get_all_models() call
|
||||
models = request.app.state.OLLAMA_MODELS
|
||||
if not models or model not in models:
|
||||
await get_all_models(request, user=user)
|
||||
models = request.app.state.OLLAMA_MODELS
|
||||
|
||||
if model in models:
|
||||
url_idx = random.choice(models[model]["urls"])
|
||||
else:
|
||||
|
||||
@@ -991,8 +991,13 @@ async def generate_chat_completion(
|
||||
detail="Model not found",
|
||||
)
|
||||
|
||||
await get_all_models(request, user=user)
|
||||
model = request.app.state.OPENAI_MODELS.get(model_id)
|
||||
# Check if model is already in app state cache to avoid expensive get_all_models() call
|
||||
# This significantly reduces TTFT when models are already cached
|
||||
model = request.app.state.OPENAI_MODELS.get(model_id) if request.app.state.OPENAI_MODELS else None
|
||||
if not model:
|
||||
await get_all_models(request, user=user)
|
||||
model = request.app.state.OPENAI_MODELS.get(model_id)
|
||||
|
||||
if model:
|
||||
idx = model["urlIdx"]
|
||||
else:
|
||||
@@ -1148,9 +1153,12 @@ async def embeddings(request: Request, form_data: dict, user):
|
||||
# Prepare payload/body
|
||||
body = json.dumps(form_data)
|
||||
# Find correct backend url/key based on model
|
||||
await get_all_models(request, user=user)
|
||||
model_id = form_data.get("model")
|
||||
# Check if model is already in app state cache to avoid expensive get_all_models() call
|
||||
models = request.app.state.OPENAI_MODELS
|
||||
if not models or model_id not in models:
|
||||
await get_all_models(request, user=user)
|
||||
models = request.app.state.OPENAI_MODELS
|
||||
if model_id in models:
|
||||
idx = models[model_id]["urlIdx"]
|
||||
|
||||
|
||||
Reference in New Issue
Block a user