|
6 | 6 | from threading import Thread |
7 | 7 | from typing import Dict, List |
8 | 8 |
|
9 | | -import requests_cache |
10 | 9 | from bson import ObjectId |
11 | 10 | from fastapi import APIRouter, Depends, Request |
12 | 11 | from square_auth.auth import Auth |
|
25 | 24 | from skill_manager.keycloak_api import KeycloakAPI |
26 | 25 | from skill_manager.models import Prediction, Skill, SkillType |
27 | 26 | from skill_manager.routers import client_credentials |
| 27 | +from skill_manager.utils import merge_dicts |
28 | 28 |
|
29 | 29 | logger = logging.getLogger(__name__) |
30 | 30 |
|
@@ -236,13 +236,29 @@ async def query_skill( |
236 | 236 | query = query_request.query |
237 | 237 | user_id = query_request.user_id |
238 | 238 |
|
239 | | - skill = await get_skill_if_authorized(request, skill_id=id, write_access=False) |
| 239 | + skill: Skill = await get_skill_if_authorized( |
| 240 | + request, skill_id=id, write_access=False |
| 241 | + ) |
240 | 242 | query_request.skill = json.loads(skill.json()) |
241 | 243 |
|
242 | | - default_skill_args = skill.default_skill_args |
243 | | - if default_skill_args is not None: |
244 | | - # add default skill args, potentially overwrite with query.skill_args |
245 | | - query_request.skill_args = {**default_skill_args, **query_request.skill_args} |
| 244 | + # merge kargs with kwargs in request |
| 245 | + for kwargs_key in [ |
| 246 | + "model_kwargs", |
| 247 | + "task_kwargs", |
| 248 | + "preprocessing_kwargs", |
| 249 | + "explain_kwargs", |
| 250 | + "attack_kwargs", |
| 251 | + ]: |
| 252 | + # overwrite kwargs from query_request with default kwargs |
| 253 | + kwargs = merge_dicts( |
| 254 | + skill.default_skill_args.pop(kwargs_key, {}), |
| 255 | + getattr(query_request, kwargs_key), |
| 256 | + ) |
| 257 | + # set kwargs in query_request |
| 258 | + setattr(query_request, kwargs_key, kwargs) |
| 259 | + query_request.skill_args = merge_dicts( |
| 260 | + skill.default_skill_args, query_request.skill_args |
| 261 | + ) |
246 | 262 |
|
247 | 263 | headers = {"Authorization": f"Bearer {token}"} |
248 | 264 | if request.headers.get("Cache-Control"): |
|
0 commit comments