This commit is contained in:
Timothy Jaeryang Baek
2026-03-25 02:49:34 -05:00
parent 08ff3bd30f
commit 857d7e6f37
2 changed files with 49 additions and 6 deletions
+20 -6
View File
@@ -203,27 +203,41 @@ async def generate_chat_completion(
except Exception as e:
raise e
if model.get('owned_by') == 'arena':
# Arena model — sub-model was already resolved by process_chat_payload.
# Inject selected_model_id into the response for the frontend.
metadata = form_data.get('metadata', {})
selected_model_id = metadata.pop('selected_model_id', None)
# Also clear from request.state.metadata to prevent the merge at
# lines 177-179 from re-adding it on the recursive call.
if hasattr(request.state, 'metadata'):
request.state.metadata.pop('selected_model_id', None)
# Fallback: if generate_chat_completion is called with an arena model
# from a path that did NOT go through process_chat_payload (e.g.,
# background tasks for title/follow-up/tags generation), resolve now.
if not selected_model_id and model.get('owned_by') == 'arena':
model_ids = model.get('info', {}).get('meta', {}).get('model_ids')
filter_mode = model.get('info', {}).get('meta', {}).get('filter_mode')
if model_ids and filter_mode == 'exclude':
model_ids = [
model['id']
for model in list(request.app.state.MODELS.values())
if model.get('owned_by') != 'arena' and model['id'] not in model_ids
available_model['id']
for available_model in list(request.app.state.MODELS.values())
if available_model.get('owned_by') != 'arena' and available_model['id'] not in model_ids
]
selected_model_id = None
if isinstance(model_ids, list) and model_ids:
selected_model_id = random.choice(model_ids)
else:
model_ids = [
model['id'] for model in list(request.app.state.MODELS.values()) if model.get('owned_by') != 'arena'
available_model['id']
for available_model in list(request.app.state.MODELS.values())
if available_model.get('owned_by') != 'arena'
]
selected_model_id = random.choice(model_ids)
form_data['model'] = selected_model_id
if selected_model_id:
if form_data.get('stream') == True:
async def stream_wrapper(stream):