Fix task generation (title, tags, etc.) failing for direct connection chats

Two fixes for direct connection users where background task generation
(titles, tags, follow-ups) crashes with TypeError or returns 404:

1. middleware.py: Clear request.state.direct before running
   background_tasks_handler so task generation routes through the
   server-side path instead of generate_direct_chat_completion, which
   requires an active WebSocket session unavailable in this context.

2. tasks.py: Always use request.app.state.MODELS for task endpoints and
   move model validation after get_task_model_id() so the configured
   task model can resolve before the existence check runs. Previously,
   the chat model (from a direct connection) was validated before the
   task model override, causing a 404 when the chat model wasn't in the
   server's MODELS registry.

Fixes #24092, Fixes #24095

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
This commit is contained in:
Tim Champ 2026-04-24 13:30:46 -04:00
parent 4e2240aada
commit d827039b1b
2 changed files with 94 additions and 97 deletions

View file

@ -157,22 +157,14 @@ async def generate_title(request: Request, form_data: dict, user=Depends(get_ver
content={'detail': 'Title generation is disabled'},
)
if getattr(request.state, 'direct', False) and hasattr(request.state, 'model'):
models = {
request.state.model['id']: request.state.model,
}
else:
models = request.app.state.MODELS
models = request.app.state.MODELS
model_id = form_data['model']
if model_id not in models:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail=ERROR_MESSAGES.MODEL_NOT_FOUND(),
)
# Check if the user has a custom task model
# If the user has a custom task model, use that model
# Resolve the task model before validating — the configured task model
# may differ from the chat model (e.g., when users chat via direct
# connections, the chat model may not be in the server's MODELS registry,
# but the configured task model is).
task_model_id = get_task_model_id(
model_id,
request.app.state.config.TASK_MODEL,
@ -180,6 +172,12 @@ async def generate_title(request: Request, form_data: dict, user=Depends(get_ver
models,
)
if task_model_id not in models:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail=ERROR_MESSAGES.MODEL_NOT_FOUND(),
)
log.debug(f'generating chat title using model {task_model_id} for user {user.email} ')
if request.app.state.config.TITLE_GENERATION_PROMPT_TEMPLATE != '':
@ -234,22 +232,14 @@ async def generate_follow_ups(request: Request, form_data: dict, user=Depends(ge
content={'detail': 'Follow-up generation is disabled'},
)
if getattr(request.state, 'direct', False) and hasattr(request.state, 'model'):
models = {
request.state.model['id']: request.state.model,
}
else:
models = request.app.state.MODELS
models = request.app.state.MODELS
model_id = form_data['model']
if model_id not in models:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail=ERROR_MESSAGES.MODEL_NOT_FOUND(),
)
# Check if the user has a custom task model
# If the user has a custom task model, use that model
# Resolve the task model before validating — the configured task model
# may differ from the chat model (e.g., when users chat via direct
# connections, the chat model may not be in the server's MODELS registry,
# but the configured task model is).
task_model_id = get_task_model_id(
model_id,
request.app.state.config.TASK_MODEL,
@ -257,6 +247,12 @@ async def generate_follow_ups(request: Request, form_data: dict, user=Depends(ge
models,
)
if task_model_id not in models:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail=ERROR_MESSAGES.MODEL_NOT_FOUND(),
)
log.debug(f'generating chat title using model {task_model_id} for user {user.email} ')
if request.app.state.config.FOLLOW_UP_GENERATION_PROMPT_TEMPLATE != '':
@ -302,22 +298,14 @@ async def generate_chat_tags(request: Request, form_data: dict, user=Depends(get
content={'detail': 'Tags generation is disabled'},
)
if getattr(request.state, 'direct', False) and hasattr(request.state, 'model'):
models = {
request.state.model['id']: request.state.model,
}
else:
models = request.app.state.MODELS
models = request.app.state.MODELS
model_id = form_data['model']
if model_id not in models:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail=ERROR_MESSAGES.MODEL_NOT_FOUND(),
)
# Check if the user has a custom task model
# If the user has a custom task model, use that model
# Resolve the task model before validating — the configured task model
# may differ from the chat model (e.g., when users chat via direct
# connections, the chat model may not be in the server's MODELS registry,
# but the configured task model is).
task_model_id = get_task_model_id(
model_id,
request.app.state.config.TASK_MODEL,
@ -325,6 +313,12 @@ async def generate_chat_tags(request: Request, form_data: dict, user=Depends(get
models,
)
if task_model_id not in models:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail=ERROR_MESSAGES.MODEL_NOT_FOUND(),
)
log.debug(f'generating chat tags using model {task_model_id} for user {user.email} ')
if request.app.state.config.TAGS_GENERATION_PROMPT_TEMPLATE != '':
@ -364,22 +358,14 @@ async def generate_chat_tags(request: Request, form_data: dict, user=Depends(get
@router.post('/image_prompt/completions')
async def generate_image_prompt(request: Request, form_data: dict, user=Depends(get_verified_user)):
if getattr(request.state, 'direct', False) and hasattr(request.state, 'model'):
models = {
request.state.model['id']: request.state.model,
}
else:
models = request.app.state.MODELS
models = request.app.state.MODELS
model_id = form_data['model']
if model_id not in models:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail=ERROR_MESSAGES.MODEL_NOT_FOUND(),
)
# Check if the user has a custom task model
# If the user has a custom task model, use that model
# Resolve the task model before validating — the configured task model
# may differ from the chat model (e.g., when users chat via direct
# connections, the chat model may not be in the server's MODELS registry,
# but the configured task model is).
task_model_id = get_task_model_id(
model_id,
request.app.state.config.TASK_MODEL,
@ -387,6 +373,12 @@ async def generate_image_prompt(request: Request, form_data: dict, user=Depends(
models,
)
if task_model_id not in models:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail=ERROR_MESSAGES.MODEL_NOT_FOUND(),
)
log.debug(f'generating image prompt using model {task_model_id} for user {user.email} ')
if request.app.state.config.IMAGE_PROMPT_GENERATION_PROMPT_TEMPLATE != '':
@ -444,22 +436,14 @@ async def generate_queries(request: Request, form_data: dict, user=Depends(get_v
log.info(f'Reusing cached queries: {request.state.cached_queries}')
return request.state.cached_queries
if getattr(request.state, 'direct', False) and hasattr(request.state, 'model'):
models = {
request.state.model['id']: request.state.model,
}
else:
models = request.app.state.MODELS
models = request.app.state.MODELS
model_id = form_data['model']
if model_id not in models:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail=ERROR_MESSAGES.MODEL_NOT_FOUND(),
)
# Check if the user has a custom task model
# If the user has a custom task model, use that model
# Resolve the task model before validating — the configured task model
# may differ from the chat model (e.g., when users chat via direct
# connections, the chat model may not be in the server's MODELS registry,
# but the configured task model is).
task_model_id = get_task_model_id(
model_id,
request.app.state.config.TASK_MODEL,
@ -467,6 +451,12 @@ async def generate_queries(request: Request, form_data: dict, user=Depends(get_v
models,
)
if task_model_id not in models:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail=ERROR_MESSAGES.MODEL_NOT_FOUND(),
)
log.debug(f'generating {type} queries using model {task_model_id} for user {user.email}')
if (request.app.state.config.QUERY_GENERATION_PROMPT_TEMPLATE).strip() != '':
@ -522,22 +512,14 @@ async def generate_autocompletion(request: Request, form_data: dict, user=Depend
detail=ERROR_MESSAGES.INPUT_TOO_LONG(request.app.state.config.AUTOCOMPLETE_GENERATION_INPUT_MAX_LENGTH),
)
if getattr(request.state, 'direct', False) and hasattr(request.state, 'model'):
models = {
request.state.model['id']: request.state.model,
}
else:
models = request.app.state.MODELS
models = request.app.state.MODELS
model_id = form_data['model']
if model_id not in models:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail=ERROR_MESSAGES.MODEL_NOT_FOUND(),
)
# Check if the user has a custom task model
# If the user has a custom task model, use that model
# Resolve the task model before validating — the configured task model
# may differ from the chat model (e.g., when users chat via direct
# connections, the chat model may not be in the server's MODELS registry,
# but the configured task model is).
task_model_id = get_task_model_id(
model_id,
request.app.state.config.TASK_MODEL,
@ -545,6 +527,12 @@ async def generate_autocompletion(request: Request, form_data: dict, user=Depend
models,
)
if task_model_id not in models:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail=ERROR_MESSAGES.MODEL_NOT_FOUND(),
)
log.debug(f'generating autocompletion using model {task_model_id} for user {user.email}')
if (request.app.state.config.AUTOCOMPLETE_GENERATION_PROMPT_TEMPLATE).strip() != '':
@ -584,22 +572,14 @@ async def generate_autocompletion(request: Request, form_data: dict, user=Depend
@router.post('/emoji/completions')
async def generate_emoji(request: Request, form_data: dict, user=Depends(get_verified_user)):
if getattr(request.state, 'direct', False) and hasattr(request.state, 'model'):
models = {
request.state.model['id']: request.state.model,
}
else:
models = request.app.state.MODELS
models = request.app.state.MODELS
model_id = form_data['model']
if model_id not in models:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail=ERROR_MESSAGES.MODEL_NOT_FOUND(),
)
# Check if the user has a custom task model
# If the user has a custom task model, use that model
# Resolve the task model before validating — the configured task model
# may differ from the chat model (e.g., when users chat via direct
# connections, the chat model may not be in the server's MODELS registry,
# but the configured task model is).
task_model_id = get_task_model_id(
model_id,
request.app.state.config.TASK_MODEL,
@ -607,6 +587,12 @@ async def generate_emoji(request: Request, form_data: dict, user=Depends(get_ver
models,
)
if task_model_id not in models:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail=ERROR_MESSAGES.MODEL_NOT_FOUND(),
)
log.debug(f'generating emoji using model {task_model_id} for user {user.email} ')
template = DEFAULT_EMOJI_GENERATION_PROMPT_TEMPLATE
@ -649,12 +635,7 @@ async def generate_emoji(request: Request, form_data: dict, user=Depends(get_ver
@router.post('/moa/completions')
async def generate_moa_response(request: Request, form_data: dict, user=Depends(get_verified_user)):
if getattr(request.state, 'direct', False) and hasattr(request.state, 'model'):
models = {
request.state.model['id']: request.state.model,
}
else:
models = request.app.state.MODELS
models = request.app.state.MODELS
model_id = form_data['model']

View file

@ -3006,6 +3006,17 @@ async def background_tasks_handler(ctx):
tasks = ctx['tasks']
event_emitter = ctx['event_emitter']
# Background task generation (titles, tags, follow-ups) should always use
# the server-side task model, not the user's direct connection. Clear the
# direct flag so generate_chat_completion routes through the standard
# server-side path instead of generate_direct_chat_completion, which
# requires an active WebSocket session that doesn't exist in this context.
original_direct = getattr(request.state, 'direct', False)
original_model = getattr(request.state, 'model', None)
request.state.direct = False
if hasattr(request.state, 'model'):
del request.state.model
message = None
messages = []
@ -3200,6 +3211,11 @@ async def background_tasks_handler(ctx):
except Exception as e:
pass
# Restore original direct connection state
request.state.direct = original_direct
if original_model is not None:
request.state.model = original_model
async def outlet_filter_handler(ctx):
"""Run outlet filters inline after chat completion.