diff --git a/app/desktop/studio_server/copilot_api.py b/app/desktop/studio_server/copilot_api.py index 0b3090499..d48f9e413 100644 --- a/app/desktop/studio_server/copilot_api.py +++ b/app/desktop/studio_server/copilot_api.py @@ -12,7 +12,6 @@ ClarifySpecOutput, GenerateBatchInput, GenerateBatchOutput, - HTTPValidationError, RefineSpecInput, ) from app.desktop.studio_server.api_client.kiln_ai_server_client.models import ( @@ -43,11 +42,11 @@ TaskInfoApi, ) from app.desktop.studio_server.utils.copilot_utils import ( - check_response_error, create_dataset_task_runs, generate_copilot_examples, get_copilot_api_key, ) +from app.desktop.studio_server.utils.response_utils import unwrap_response from fastapi import FastAPI, HTTPException from kiln_ai.datamodel import TaskRun from kiln_ai.datamodel.basemodel import FilenameString @@ -127,19 +126,10 @@ async def clarify_spec(input: ClarifySpecApiInput) -> ClarifySpecApiOutput: body=clarify_input, ) ) - check_response_error(detailed_result) - - result = detailed_result.parsed - if result is None: - raise HTTPException( - status_code=500, detail="Failed to analyze spec. Please try again." - ) - - if isinstance(result, HTTPValidationError): - raise HTTPException( - status_code=422, - detail="Validation error.", - ) + result = unwrap_response( + detailed_result, + none_detail="Failed to analyze spec. Please try again.", + ) if isinstance(result, ClarifySpecOutput): return ClarifySpecApiOutput.model_validate(result.to_dict()) @@ -162,20 +152,10 @@ async def refine_spec(input: RefineSpecApiInput) -> RefineSpecApiOutput: body=refine_input, ) ) - check_response_error(detailed_result) - - result = detailed_result.parsed - if result is None: - raise HTTPException( - status_code=500, - detail="Failed to refine spec with feedback. Please try again.", - ) - - if isinstance(result, HTTPValidationError): - raise HTTPException( - status_code=422, - detail="Validation error.", - ) + result = unwrap_response( + detailed_result, + none_detail="Failed to refine spec with feedback. Please try again.", + ) if isinstance(result, RefineSpecApiOutputClient): return RefineSpecApiOutput.model_validate(result.to_dict()) @@ -198,20 +178,10 @@ async def generate_batch(input: GenerateBatchApiInput) -> GenerateBatchApiOutput body=generate_input, ) ) - check_response_error(detailed_result) - - result = detailed_result.parsed - if result is None: - raise HTTPException( - status_code=500, - detail="Failed to generate synthetic data for spec. Please try again.", - ) - - if isinstance(result, HTTPValidationError): - raise HTTPException( - status_code=422, - detail="Validation error.", - ) + result = unwrap_response( + detailed_result, + none_detail="Failed to generate synthetic data for spec. Please try again.", + ) if isinstance(result, GenerateBatchOutput): return GenerateBatchApiOutput.model_validate(result.to_dict()) @@ -236,20 +206,10 @@ async def question_spec( body=questioner_input, ) ) - check_response_error(detailed_result) - - result = detailed_result.parsed - if result is None: - raise HTTPException( - status_code=500, - detail="Failed to generate clarifying questions for spec. Please try again.", - ) - - if isinstance(result, HTTPValidationError): - raise HTTPException( - status_code=422, - detail="Validation error.", - ) + result = unwrap_response( + detailed_result, + none_detail="Failed to generate clarifying questions for spec. Please try again.", + ) if isinstance(result, QuestionSetServerApi): return QuestionSet.model_validate(result.to_dict()) @@ -272,20 +232,10 @@ async def submit_question_answers( client=client, body=submit_input, ) - check_response_error(detailed_result) - - result = detailed_result.parsed - if result is None: - raise HTTPException( - status_code=500, - detail="Failed to refine spec with question answers. Please try again.", - ) - - if isinstance(result, HTTPValidationError): - raise HTTPException( - status_code=422, - detail="Validation error.", - ) + result = unwrap_response( + detailed_result, + none_detail="Failed to refine spec with question answers. Please try again.", + ) if isinstance(result, RefineSpecApiOutputClient): return RefineSpecApiOutput.model_validate(result.to_dict()) diff --git a/app/desktop/studio_server/prompt_optimization_job_api.py b/app/desktop/studio_server/prompt_optimization_job_api.py index 5bd039d6f..c44ecc78e 100644 --- a/app/desktop/studio_server/prompt_optimization_job_api.py +++ b/app/desktop/studio_server/prompt_optimization_job_api.py @@ -16,9 +16,6 @@ from app.desktop.studio_server.api_client.kiln_ai_server_client.models.body_start_prompt_optimization_job_v1_jobs_prompt_optimization_job_start_post import ( BodyStartPromptOptimizationJobV1JobsPromptOptimizationJobStartPost, ) -from app.desktop.studio_server.api_client.kiln_ai_server_client.models.http_validation_error import ( - HTTPValidationError, -) from app.desktop.studio_server.api_client.kiln_ai_server_client.models.job_status import ( JobStatus, ) @@ -34,7 +31,7 @@ eval_from_id, task_run_config_from_id, ) -from app.desktop.studio_server.utils.copilot_utils import check_response_error +from app.desktop.studio_server.utils.response_utils import unwrap_response from fastapi import FastAPI, HTTPException from kiln_ai.cli.commands.package_project import ( PackageForTrainingConfig, @@ -240,17 +237,16 @@ async def _create_artifacts_for_succeeded_job( prompt_optimization_job.optimized_prompt = reloaded_job.optimized_prompt return - result_response = await get_prompt_optimization_job_result_v1_jobs_prompt_optimization_job_job_id_result_get.asyncio( + detailed_response = await get_prompt_optimization_job_result_v1_jobs_prompt_optimization_job_job_id_result_get.asyncio_detailed( job_id=prompt_optimization_job.job_id, client=server_client, ) + result_response = unwrap_response( + detailed_response, + default_detail="Failed to get Prompt Optimization job result.", + ) - if ( - result_response - and not isinstance(result_response, HTTPValidationError) - and result_response.output - and hasattr(result_response.output, "optimized_prompt") - ): + if result_response.output and result_response.output.optimized_prompt: optimized_prompt_text = result_response.output.optimized_prompt prompt_optimization_job.optimized_prompt = optimized_prompt_text @@ -295,21 +291,17 @@ async def update_prompt_optimization_job_and_create_artifacts( ) try: - status_response = ( - await get_job_status_v1_jobs_job_type_job_id_status_get.asyncio( + detailed_response = ( + await get_job_status_v1_jobs_job_type_job_id_status_get.asyncio_detailed( job_type=JobType.GEPA_JOB, job_id=prompt_optimization_job.job_id, client=server_client, ) ) - - if status_response is None or isinstance(status_response, HTTPValidationError): - logger.warning( - f"Could not fetch status for Prompt Optimization job {prompt_optimization_job.job_id}" - ) - raise RuntimeError( - f"Could not fetch status for Prompt Optimization job {prompt_optimization_job.job_id}: {status_response}" - ) + status_response = unwrap_response( + detailed_response, + default_detail=f"Could not fetch status for Prompt Optimization job {prompt_optimization_job.job_id}", + ) new_status = str(status_response.status.value) @@ -370,25 +362,12 @@ async def check_run_config( status_code=500, detail="Server client not authenticated" ) - response = await check_prompt_optimization_model_supported_v1_jobs_prompt_optimization_job_check_model_supported_get.asyncio( + detailed_response = await check_prompt_optimization_model_supported_v1_jobs_prompt_optimization_job_check_model_supported_get.asyncio_detailed( client=server_client, model_name=model_name, model_provider_name=model_provider.value, ) - - if isinstance(response, HTTPValidationError): - error_detail = ( - str(response.detail) - if hasattr(response, "detail") - else "Validation error" - ) - raise HTTPException(status_code=422, detail=error_detail) - - if response is None: - raise HTTPException( - status_code=500, - detail="Failed to check run config: No response from server", - ) + response = unwrap_response(detailed_response) return CheckRunConfigResponse(is_supported=response.is_model_supported) @@ -451,25 +430,12 @@ async def check_eval( ) # EvalConfig.model_provider is already a string, no need for .value - response = await check_prompt_optimization_model_supported_v1_jobs_prompt_optimization_job_check_model_supported_get.asyncio( + detailed_response = await check_prompt_optimization_model_supported_v1_jobs_prompt_optimization_job_check_model_supported_get.asyncio_detailed( client=server_client, model_name=model_name, model_provider_name=model_provider, ) - - if isinstance(response, HTTPValidationError): - error_detail = ( - str(response.detail) - if hasattr(response, "detail") - else "Validation error" - ) - raise HTTPException(status_code=422, detail=error_detail) - - if response is None: - raise HTTPException( - status_code=500, - detail="Failed to check eval: No response from server", - ) + response = unwrap_response(detailed_response) return CheckEvalResponse( has_default_config=True, @@ -560,17 +526,7 @@ async def start_prompt_optimization_job( detailed_response = await start_prompt_optimization_job_v1_jobs_prompt_optimization_job_start_post.asyncio_detailed( client=server_client, body=body ) - check_response_error( - detailed_response, - default_detail="Failed to start Prompt Optimization job: unexpected error from server", - ) - - response = detailed_response.parsed - if response is None or isinstance(response, HTTPValidationError): - raise HTTPException( - status_code=500, - detail="Failed to start Prompt Optimization job: unexpected response from server", - ) + response = unwrap_response(detailed_response) prompt_optimization_job = PromptOptimizationJob( name=generate_memorable_name(), @@ -695,17 +651,15 @@ async def get_prompt_optimization_job_status( status_code=500, detail="Server client not authenticated" ) - response = await get_job_status_v1_jobs_job_type_job_id_status_get.asyncio( + detailed_response = await get_job_status_v1_jobs_job_type_job_id_status_get.asyncio_detailed( job_type=JobType.GEPA_JOB, job_id=job_id, client=server_client, ) - - if response is None or isinstance(response, HTTPValidationError): - raise HTTPException( - status_code=404, - detail=f"Prompt Optimization job {job_id} not found", - ) + response = unwrap_response( + detailed_response, + default_detail=f"Prompt Optimization job {job_id} not found", + ) return PublicPromptOptimizationJobStatusResponse( job_id=response.job_id, status=response.status @@ -736,16 +690,14 @@ async def get_prompt_optimization_job_result( status_code=500, detail="Server client not authenticated" ) - response = await get_prompt_optimization_job_result_v1_jobs_prompt_optimization_job_job_id_result_get.asyncio( + detailed_response = await get_prompt_optimization_job_result_v1_jobs_prompt_optimization_job_job_id_result_get.asyncio_detailed( job_id=job_id, client=server_client, ) - - if response is None or isinstance(response, HTTPValidationError): - raise HTTPException( - status_code=404, - detail=f"Prompt Optimization job {job_id} result not found", - ) + response = unwrap_response( + detailed_response, + default_detail=f"Prompt Optimization job {job_id} result not found", + ) if not response.output or not hasattr(response.output, "optimized_prompt"): raise HTTPException( diff --git a/app/desktop/studio_server/settings_api.py b/app/desktop/studio_server/settings_api.py index f015226cc..a5662ce9d 100644 --- a/app/desktop/studio_server/settings_api.py +++ b/app/desktop/studio_server/settings_api.py @@ -7,10 +7,8 @@ from app.desktop.studio_server.api_client.kiln_server_client import ( get_authenticated_client, ) -from app.desktop.studio_server.utils.copilot_utils import ( - check_response_error, - get_copilot_api_key, -) +from app.desktop.studio_server.utils.copilot_utils import get_copilot_api_key +from app.desktop.studio_server.utils.response_utils import unwrap_response from fastapi import FastAPI, HTTPException from kiln_ai.utils.config import Config from kiln_ai.utils.filesystem import open_folder @@ -79,13 +77,9 @@ async def check_entitlements(feature_codes: str) -> dict[str, bool]: feature_codes=feature_codes, ) ) - check_response_error(detailed_result) - - result = detailed_result.parsed - if result is None: - raise HTTPException( - status_code=500, - detail="Failed to check entitlements. Please try again.", - ) + result = unwrap_response( + detailed_result, + none_detail="Failed to check entitlements. Please try again.", + ) return result.additional_properties diff --git a/app/desktop/studio_server/test_prompt_optimization_job_api.py b/app/desktop/studio_server/test_prompt_optimization_job_api.py index b0a935fd5..ddac27221 100644 --- a/app/desktop/studio_server/test_prompt_optimization_job_api.py +++ b/app/desktop/studio_server/test_prompt_optimization_job_api.py @@ -8,9 +8,6 @@ from app.desktop.studio_server.api_client.kiln_ai_server_client.client import ( AuthenticatedClient, ) -from app.desktop.studio_server.api_client.kiln_ai_server_client.models.http_validation_error import ( - HTTPValidationError, -) from app.desktop.studio_server.api_client.kiln_ai_server_client.models.job_start_response import ( JobStartResponse, ) @@ -109,9 +106,9 @@ def test_get_prompt_optimization_job_result_success(client, mock_api_key): ) with patch( - "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_prompt_optimization_job_result_v1_jobs_prompt_optimization_job_job_id_result_get.asyncio", + "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_prompt_optimization_job_result_v1_jobs_prompt_optimization_job_job_id_result_get.asyncio_detailed", new_callable=AsyncMock, - return_value=mock_response, + return_value=_make_sdk_response(parsed=mock_response), ): response = client.get(f"/api/prompt_optimization_jobs/{job_id}/result") @@ -127,17 +124,16 @@ def test_get_prompt_optimization_job_result_not_found(client, mock_api_key): job_id = "nonexistent-job" with patch( - "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_prompt_optimization_job_result_v1_jobs_prompt_optimization_job_job_id_result_get.asyncio", + "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_prompt_optimization_job_result_v1_jobs_prompt_optimization_job_job_id_result_get.asyncio_detailed", new_callable=AsyncMock, - return_value=None, + return_value=_make_sdk_response( + status_code=HTTPStatus.NOT_FOUND, + content=b'{"message": "Not found"}', + ), ): response = client.get(f"/api/prompt_optimization_jobs/{job_id}/result") assert response.status_code == 404 - assert ( - f"Prompt Optimization job {job_id} result not found" - in response.json()["detail"] - ) def test_get_prompt_optimization_job_result_no_output(client, mock_api_key): @@ -149,9 +145,9 @@ def test_get_prompt_optimization_job_result_no_output(client, mock_api_key): ) with patch( - "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_prompt_optimization_job_result_v1_jobs_prompt_optimization_job_job_id_result_get.asyncio", + "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_prompt_optimization_job_result_v1_jobs_prompt_optimization_job_job_id_result_get.asyncio_detailed", new_callable=AsyncMock, - return_value=mock_response, + return_value=_make_sdk_response(parsed=mock_response), ): response = client.get(f"/api/prompt_optimization_jobs/{job_id}/result") @@ -164,7 +160,7 @@ def test_get_prompt_optimization_job_result_api_error(client, mock_api_key): job_id = "test-job-error" with patch( - "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_prompt_optimization_job_result_v1_jobs_prompt_optimization_job_job_id_result_get.asyncio", + "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_prompt_optimization_job_result_v1_jobs_prompt_optimization_job_job_id_result_get.asyncio_detailed", new_callable=AsyncMock, side_effect=Exception("API connection failed"), ): @@ -183,9 +179,9 @@ def test_get_prompt_optimization_job_status_success(client, mock_api_key): mock_response = JobStatusResponse(job_id=job_id, status=JobStatus.RUNNING) with patch( - "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_job_status_v1_jobs_job_type_job_id_status_get.asyncio", + "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_job_status_v1_jobs_job_type_job_id_status_get.asyncio_detailed", new_callable=AsyncMock, - return_value=mock_response, + return_value=_make_sdk_response(parsed=mock_response), ): response = client.get(f"/api/prompt_optimization_jobs/{job_id}/status") @@ -202,16 +198,16 @@ def test_get_prompt_optimization_job_status_not_found(client, mock_api_key): job_id = "nonexistent-job" with patch( - "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_job_status_v1_jobs_job_type_job_id_status_get.asyncio", + "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_job_status_v1_jobs_job_type_job_id_status_get.asyncio_detailed", new_callable=AsyncMock, - return_value=None, + return_value=_make_sdk_response( + status_code=HTTPStatus.NOT_FOUND, + content=b'{"message": "Not found"}', + ), ): response = client.get(f"/api/prompt_optimization_jobs/{job_id}/status") assert response.status_code == 404 - assert ( - f"Prompt Optimization job {job_id} not found" in response.json()["detail"] - ) def test_get_prompt_optimization_job_result_no_api_key(client): @@ -234,16 +230,17 @@ def test_get_prompt_optimization_job_result_validation_error(client, mock_api_ke """Test getting Prompt Optimization job result with validation error from server.""" job_id = "test-job-123" - mock_error = HTTPValidationError(detail=[]) - with patch( - "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_prompt_optimization_job_result_v1_jobs_prompt_optimization_job_job_id_result_get.asyncio", + "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_prompt_optimization_job_result_v1_jobs_prompt_optimization_job_job_id_result_get.asyncio_detailed", new_callable=AsyncMock, - return_value=mock_error, + return_value=_make_sdk_response( + status_code=HTTPStatus.UNPROCESSABLE_ENTITY, + content=b'{"message": "Validation error"}', + ), ): response = client.get(f"/api/prompt_optimization_jobs/{job_id}/result") - assert response.status_code == 404 + assert response.status_code == 422 def test_public_prompt_optimization_job_result_response_model(): @@ -500,9 +497,9 @@ def test_get_prompt_optimization_job_detail(client, mock_api_key, tmp_path): return_value=task, ), patch( - "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_job_status_v1_jobs_job_type_job_id_status_get.asyncio", + "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_job_status_v1_jobs_job_type_job_id_status_get.asyncio_detailed", new_callable=AsyncMock, - return_value=mock_status_response, + return_value=_make_sdk_response(parsed=mock_status_response), ), ): response = client.get( @@ -571,14 +568,14 @@ def test_prompt_optimization_job_creates_prompt_on_success( return_value=task, ), patch( - "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_job_status_v1_jobs_job_type_job_id_status_get.asyncio", + "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_job_status_v1_jobs_job_type_job_id_status_get.asyncio_detailed", new_callable=AsyncMock, - return_value=mock_status_response, + return_value=_make_sdk_response(parsed=mock_status_response), ), patch( - "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_prompt_optimization_job_result_v1_jobs_prompt_optimization_job_job_id_result_get.asyncio", + "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_prompt_optimization_job_result_v1_jobs_prompt_optimization_job_job_id_result_get.asyncio_detailed", new_callable=AsyncMock, - return_value=mock_result_response, + return_value=_make_sdk_response(parsed=mock_result_response), ), patch( "app.desktop.studio_server.prompt_optimization_job_api.task_run_config_from_id", @@ -656,14 +653,14 @@ def test_prompt_optimization_job_only_creates_prompt_once( return_value=task, ), patch( - "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_job_status_v1_jobs_job_type_job_id_status_get.asyncio", + "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_job_status_v1_jobs_job_type_job_id_status_get.asyncio_detailed", new_callable=AsyncMock, - return_value=mock_status_response, + return_value=_make_sdk_response(parsed=mock_status_response), ), patch( - "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_prompt_optimization_job_result_v1_jobs_prompt_optimization_job_job_id_result_get.asyncio", + "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_prompt_optimization_job_result_v1_jobs_prompt_optimization_job_job_id_result_get.asyncio_detailed", new_callable=AsyncMock, - return_value=mock_result_response, + return_value=_make_sdk_response(parsed=mock_result_response), ), patch( "app.desktop.studio_server.prompt_optimization_job_api.task_run_config_from_id", @@ -719,11 +716,11 @@ def test_get_prompt_optimization_job_skips_update_when_succeeded( return_value=task, ), patch( - "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_job_status_v1_jobs_job_type_job_id_status_get.asyncio", + "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_job_status_v1_jobs_job_type_job_id_status_get.asyncio_detailed", new_callable=AsyncMock, ) as mock_status, patch( - "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_prompt_optimization_job_result_v1_jobs_prompt_optimization_job_job_id_result_get.asyncio", + "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_prompt_optimization_job_result_v1_jobs_prompt_optimization_job_job_id_result_get.asyncio_detailed", new_callable=AsyncMock, ) as mock_result, ): @@ -774,11 +771,11 @@ def test_get_prompt_optimization_job_skips_update_when_failed( return_value=task, ), patch( - "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_job_status_v1_jobs_job_type_job_id_status_get.asyncio", + "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_job_status_v1_jobs_job_type_job_id_status_get.asyncio_detailed", new_callable=AsyncMock, ) as mock_status, patch( - "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_prompt_optimization_job_result_v1_jobs_prompt_optimization_job_job_id_result_get.asyncio", + "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_prompt_optimization_job_result_v1_jobs_prompt_optimization_job_job_id_result_get.asyncio_detailed", new_callable=AsyncMock, ) as mock_result, ): @@ -828,11 +825,11 @@ def test_get_prompt_optimization_job_skips_update_when_cancelled( return_value=task, ), patch( - "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_job_status_v1_jobs_job_type_job_id_status_get.asyncio", + "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_job_status_v1_jobs_job_type_job_id_status_get.asyncio_detailed", new_callable=AsyncMock, ) as mock_status, patch( - "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_prompt_optimization_job_result_v1_jobs_prompt_optimization_job_job_id_result_get.asyncio", + "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_prompt_optimization_job_result_v1_jobs_prompt_optimization_job_job_id_result_get.asyncio_detailed", new_callable=AsyncMock, ) as mock_result, ): @@ -1094,9 +1091,12 @@ def test_update_prompt_optimization_job_status_response_none( return_value=task, ), patch( - "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_job_status_v1_jobs_job_type_job_id_status_get.asyncio", + "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_job_status_v1_jobs_job_type_job_id_status_get.asyncio_detailed", new_callable=AsyncMock, - return_value=None, + return_value=_make_sdk_response( + status_code=HTTPStatus.INTERNAL_SERVER_ERROR, + content=b'{"message": "Server error"}', + ), ), ): response = client.get( @@ -1136,17 +1136,18 @@ def test_update_prompt_optimization_job_status_response_validation_error( task_id = task.id prompt_optimization_job_id = prompt_optimization_job.id - mock_error = HTTPValidationError(detail=[]) - with ( patch( "app.desktop.studio_server.prompt_optimization_job_api.task_from_id", return_value=task, ), patch( - "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_job_status_v1_jobs_job_type_job_id_status_get.asyncio", + "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_job_status_v1_jobs_job_type_job_id_status_get.asyncio_detailed", new_callable=AsyncMock, - return_value=mock_error, + return_value=_make_sdk_response( + status_code=HTTPStatus.UNPROCESSABLE_ENTITY, + content=b'{"message": "Validation error"}', + ), ), ): response = client.get( @@ -1192,7 +1193,7 @@ def test_update_prompt_optimization_job_status_exception_during_update( return_value=task, ), patch( - "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_job_status_v1_jobs_job_type_job_id_status_get.asyncio", + "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_job_status_v1_jobs_job_type_job_id_status_get.asyncio_detailed", new_callable=AsyncMock, side_effect=Exception("Network error"), ), @@ -1382,17 +1383,18 @@ def test_check_run_config_server_validation_error(client, mock_api_key, tmp_path task_id = task.id run_config_id = "test-config-id" - mock_error = HTTPValidationError(detail="Invalid model") - with ( patch( "app.desktop.studio_server.prompt_optimization_job_api.task_run_config_from_id", return_value=mock_run_config, ), patch( - "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.check_prompt_optimization_model_supported_v1_jobs_prompt_optimization_job_check_model_supported_get.asyncio", + "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.check_prompt_optimization_model_supported_v1_jobs_prompt_optimization_job_check_model_supported_get.asyncio_detailed", new_callable=AsyncMock, - return_value=mock_error, + return_value=_make_sdk_response( + status_code=HTTPStatus.UNPROCESSABLE_ENTITY, + content=b'{"message": "Invalid model"}', + ), ), ): response = client.get( @@ -1431,9 +1433,9 @@ def test_check_run_config_server_none_response(client, mock_api_key, tmp_path): return_value=mock_run_config, ), patch( - "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.check_prompt_optimization_model_supported_v1_jobs_prompt_optimization_job_check_model_supported_get.asyncio", + "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.check_prompt_optimization_model_supported_v1_jobs_prompt_optimization_job_check_model_supported_get.asyncio_detailed", new_callable=AsyncMock, - return_value=None, + return_value=_make_sdk_response(parsed=None), ), ): response = client.get( @@ -1442,7 +1444,7 @@ def test_check_run_config_server_none_response(client, mock_api_key, tmp_path): ) assert response.status_code == 500 - assert "No response from server" in response.json()["detail"] + assert "unknown error" in response.json()["detail"].lower() def test_check_run_config_exception(client, mock_api_key, tmp_path): @@ -1669,8 +1671,6 @@ def test_check_eval_server_validation_error(client, mock_api_key, tmp_path): task_id = task.id eval_id = "test-eval-id" - mock_error = HTTPValidationError(detail="Invalid model") - with ( patch( "app.desktop.studio_server.prompt_optimization_job_api.eval_from_id", @@ -1681,9 +1681,12 @@ def test_check_eval_server_validation_error(client, mock_api_key, tmp_path): return_value=mock_config, ), patch( - "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.check_prompt_optimization_model_supported_v1_jobs_prompt_optimization_job_check_model_supported_get.asyncio", + "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.check_prompt_optimization_model_supported_v1_jobs_prompt_optimization_job_check_model_supported_get.asyncio_detailed", new_callable=AsyncMock, - return_value=mock_error, + return_value=_make_sdk_response( + status_code=HTTPStatus.UNPROCESSABLE_ENTITY, + content=b'{"message": "Invalid model"}', + ), ), ): response = client.get( @@ -1728,9 +1731,9 @@ def test_check_eval_server_none_response(client, mock_api_key, tmp_path): return_value=mock_config, ), patch( - "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.check_prompt_optimization_model_supported_v1_jobs_prompt_optimization_job_check_model_supported_get.asyncio", + "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.check_prompt_optimization_model_supported_v1_jobs_prompt_optimization_job_check_model_supported_get.asyncio_detailed", new_callable=AsyncMock, - return_value=None, + return_value=_make_sdk_response(parsed=None), ), ): response = client.get( @@ -1739,7 +1742,7 @@ def test_check_eval_server_none_response(client, mock_api_key, tmp_path): ) assert response.status_code == 500 - assert "No response from server" in response.json()["detail"] + assert "unknown error" in response.json()["detail"].lower() def test_check_eval_exception(client, mock_api_key, tmp_path): @@ -1814,9 +1817,9 @@ def test_check_eval_success_train_set( return_value=mock_config, ), patch( - "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.check_prompt_optimization_model_supported_v1_jobs_prompt_optimization_job_check_model_supported_get.asyncio", + "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.check_prompt_optimization_model_supported_v1_jobs_prompt_optimization_job_check_model_supported_get.asyncio_detailed", new_callable=AsyncMock, - return_value=mock_check_response, + return_value=_make_sdk_response(parsed=mock_check_response), ), ): response = client.get( @@ -2051,7 +2054,7 @@ def test_start_prompt_optimization_job_server_none_response( ) assert response.status_code == 500 - assert "unexpected response from server" in response.json()["detail"] + assert "unknown error" in response.json()["detail"].lower() def test_start_prompt_optimization_job_connection_error(client, mock_api_key, tmp_path): @@ -2264,14 +2267,14 @@ def test_prompt_optimization_job_creates_run_config_on_success( return_value=target_run_config, ), patch( - "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_job_status_v1_jobs_job_type_job_id_status_get.asyncio", + "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_job_status_v1_jobs_job_type_job_id_status_get.asyncio_detailed", new_callable=AsyncMock, - return_value=mock_status_response, + return_value=_make_sdk_response(parsed=mock_status_response), ), patch( - "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_prompt_optimization_job_result_v1_jobs_prompt_optimization_job_job_id_result_get.asyncio", + "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_prompt_optimization_job_result_v1_jobs_prompt_optimization_job_job_id_result_get.asyncio_detailed", new_callable=AsyncMock, - return_value=mock_result_response, + return_value=_make_sdk_response(parsed=mock_result_response), ), ): response = client.get( @@ -2376,14 +2379,14 @@ def test_prompt_optimization_job_only_creates_run_config_once( return_value=target_run_config, ), patch( - "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_job_status_v1_jobs_job_type_job_id_status_get.asyncio", + "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_job_status_v1_jobs_job_type_job_id_status_get.asyncio_detailed", new_callable=AsyncMock, - return_value=mock_status_response, + return_value=_make_sdk_response(parsed=mock_status_response), ), patch( - "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_prompt_optimization_job_result_v1_jobs_prompt_optimization_job_job_id_result_get.asyncio", + "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_prompt_optimization_job_result_v1_jobs_prompt_optimization_job_job_id_result_get.asyncio_detailed", new_callable=AsyncMock, - return_value=mock_result_response, + return_value=_make_sdk_response(parsed=mock_result_response), ), ): response_1 = client.get( @@ -2447,14 +2450,14 @@ def test_prompt_optimization_job_run_config_handles_missing_target_config( return_value=task, ), patch( - "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_job_status_v1_jobs_job_type_job_id_status_get.asyncio", + "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_job_status_v1_jobs_job_type_job_id_status_get.asyncio_detailed", new_callable=AsyncMock, - return_value=mock_status_response, + return_value=_make_sdk_response(parsed=mock_status_response), ), patch( - "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_prompt_optimization_job_result_v1_jobs_prompt_optimization_job_job_id_result_get.asyncio", + "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_prompt_optimization_job_result_v1_jobs_prompt_optimization_job_job_id_result_get.asyncio_detailed", new_callable=AsyncMock, - return_value=mock_result_response, + return_value=_make_sdk_response(parsed=mock_result_response), ), ): response = client.get( @@ -2642,14 +2645,14 @@ def test_prompt_optimization_job_cleanup_prompt_when_prompt_creation_fails( return_value=task, ), patch( - "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_job_status_v1_jobs_job_type_job_id_status_get.asyncio", + "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_job_status_v1_jobs_job_type_job_id_status_get.asyncio_detailed", new_callable=AsyncMock, - return_value=mock_status_response, + return_value=_make_sdk_response(parsed=mock_status_response), ), patch( - "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_prompt_optimization_job_result_v1_jobs_prompt_optimization_job_job_id_result_get.asyncio", + "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_prompt_optimization_job_result_v1_jobs_prompt_optimization_job_job_id_result_get.asyncio_detailed", new_callable=AsyncMock, - return_value=mock_result_response, + return_value=_make_sdk_response(parsed=mock_result_response), ), patch( "app.desktop.studio_server.prompt_optimization_job_api.create_prompt_from_optimization", @@ -2725,14 +2728,14 @@ def test_prompt_optimization_job_cleanup_both_artifacts_when_run_config_fails_af with ( patch( - "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_job_status_v1_jobs_job_type_job_id_status_get.asyncio", + "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_job_status_v1_jobs_job_type_job_id_status_get.asyncio_detailed", new_callable=AsyncMock, - return_value=mock_status_response, + return_value=_make_sdk_response(parsed=mock_status_response), ), patch( - "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_prompt_optimization_job_result_v1_jobs_prompt_optimization_job_job_id_result_get.asyncio", + "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_prompt_optimization_job_result_v1_jobs_prompt_optimization_job_job_id_result_get.asyncio_detailed", new_callable=AsyncMock, - return_value=mock_result_response, + return_value=_make_sdk_response(parsed=mock_result_response), ), patch( "app.desktop.studio_server.prompt_optimization_job_api.task_run_config_from_id", @@ -2816,14 +2819,14 @@ def test_prompt_optimization_job_retry_after_cleanup(mock_api_key, tmp_path): # First attempt: run config creation fails with ( patch( - "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_job_status_v1_jobs_job_type_job_id_status_get.asyncio", + "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_job_status_v1_jobs_job_type_job_id_status_get.asyncio_detailed", new_callable=AsyncMock, - return_value=mock_status_response, + return_value=_make_sdk_response(parsed=mock_status_response), ), patch( - "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_prompt_optimization_job_result_v1_jobs_prompt_optimization_job_job_id_result_get.asyncio", + "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_prompt_optimization_job_result_v1_jobs_prompt_optimization_job_job_id_result_get.asyncio_detailed", new_callable=AsyncMock, - return_value=mock_result_response, + return_value=_make_sdk_response(parsed=mock_result_response), ), patch( "app.desktop.studio_server.prompt_optimization_job_api.task_run_config_from_id", @@ -2855,14 +2858,14 @@ def test_prompt_optimization_job_retry_after_cleanup(mock_api_key, tmp_path): # Second attempt: should succeed because cleanup cleared the IDs with ( patch( - "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_job_status_v1_jobs_job_type_job_id_status_get.asyncio", + "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_job_status_v1_jobs_job_type_job_id_status_get.asyncio_detailed", new_callable=AsyncMock, - return_value=mock_status_response, + return_value=_make_sdk_response(parsed=mock_status_response), ), patch( - "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_prompt_optimization_job_result_v1_jobs_prompt_optimization_job_job_id_result_get.asyncio", + "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_prompt_optimization_job_result_v1_jobs_prompt_optimization_job_job_id_result_get.asyncio_detailed", new_callable=AsyncMock, - return_value=mock_result_response, + return_value=_make_sdk_response(parsed=mock_result_response), ), patch( "app.desktop.studio_server.prompt_optimization_job_api.task_run_config_from_id", @@ -2942,14 +2945,14 @@ async def concurrent_updates(): """Run two concurrent update calls to test locking.""" with ( patch( - "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_job_status_v1_jobs_job_type_job_id_status_get.asyncio", + "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_job_status_v1_jobs_job_type_job_id_status_get.asyncio_detailed", new_callable=AsyncMock, - return_value=mock_status_response, + return_value=_make_sdk_response(parsed=mock_status_response), ), patch( - "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_prompt_optimization_job_result_v1_jobs_prompt_optimization_job_job_id_result_get.asyncio", + "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_prompt_optimization_job_result_v1_jobs_prompt_optimization_job_job_id_result_get.asyncio_detailed", new_callable=AsyncMock, - return_value=mock_result_response, + return_value=_make_sdk_response(parsed=mock_result_response), ), patch( "app.desktop.studio_server.prompt_optimization_job_api.task_run_config_from_id", @@ -3051,14 +3054,14 @@ def test_update_prompt_optimization_job_status_transitions( with ( patch( - "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_job_status_v1_jobs_job_type_job_id_status_get.asyncio", + "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_job_status_v1_jobs_job_type_job_id_status_get.asyncio_detailed", new_callable=AsyncMock, - return_value=mock_status_response, + return_value=_make_sdk_response(parsed=mock_status_response), ), patch( - "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_prompt_optimization_job_result_v1_jobs_prompt_optimization_job_job_id_result_get.asyncio", + "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_prompt_optimization_job_result_v1_jobs_prompt_optimization_job_job_id_result_get.asyncio_detailed", new_callable=AsyncMock, - return_value=mock_result_response, + return_value=_make_sdk_response(parsed=mock_result_response), ), patch( "app.desktop.studio_server.prompt_optimization_job_api.task_run_config_from_id", @@ -3156,14 +3159,14 @@ def test_update_prompt_optimization_job_running_to_succeeded_creates_artifacts( with ( patch( - "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_job_status_v1_jobs_job_type_job_id_status_get.asyncio", + "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_job_status_v1_jobs_job_type_job_id_status_get.asyncio_detailed", new_callable=AsyncMock, - return_value=mock_status_response, + return_value=_make_sdk_response(parsed=mock_status_response), ), patch( - "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_prompt_optimization_job_result_v1_jobs_prompt_optimization_job_job_id_result_get.asyncio", + "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_prompt_optimization_job_result_v1_jobs_prompt_optimization_job_job_id_result_get.asyncio_detailed", new_callable=AsyncMock, - return_value=mock_result_response, + return_value=_make_sdk_response(parsed=mock_result_response), ), patch( "app.desktop.studio_server.prompt_optimization_job_api.task_run_config_from_id", @@ -3228,12 +3231,12 @@ def test_update_prompt_optimization_job_succeeded_to_succeeded_no_artifacts( with ( patch( - "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_job_status_v1_jobs_job_type_job_id_status_get.asyncio", + "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_job_status_v1_jobs_job_type_job_id_status_get.asyncio_detailed", new_callable=AsyncMock, - return_value=mock_status_response, + return_value=_make_sdk_response(parsed=mock_status_response), ), patch( - "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_prompt_optimization_job_result_v1_jobs_prompt_optimization_job_job_id_result_get.asyncio", + "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_prompt_optimization_job_result_v1_jobs_prompt_optimization_job_job_id_result_get.asyncio_detailed", new_callable=AsyncMock, ) as mock_result, ): @@ -3292,12 +3295,12 @@ def test_update_prompt_optimization_job_pending_to_running_no_artifacts( with ( patch( - "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_job_status_v1_jobs_job_type_job_id_status_get.asyncio", + "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_job_status_v1_jobs_job_type_job_id_status_get.asyncio_detailed", new_callable=AsyncMock, - return_value=mock_status_response, + return_value=_make_sdk_response(parsed=mock_status_response), ), patch( - "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_prompt_optimization_job_result_v1_jobs_prompt_optimization_job_job_id_result_get.asyncio", + "app.desktop.studio_server.api_client.kiln_ai_server_client.api.jobs.get_prompt_optimization_job_result_v1_jobs_prompt_optimization_job_job_id_result_get.asyncio_detailed", new_callable=AsyncMock, ) as mock_result, ): diff --git a/app/desktop/studio_server/utils/copilot_utils.py b/app/desktop/studio_server/utils/copilot_utils.py index fd9868c30..8ae58c5c6 100644 --- a/app/desktop/studio_server/utils/copilot_utils.py +++ b/app/desktop/studio_server/utils/copilot_utils.py @@ -5,7 +5,6 @@ spec creation workflow. """ -import json import random from app.desktop.studio_server.api_client.kiln_ai_server_client.api.copilot import ( @@ -14,9 +13,7 @@ from app.desktop.studio_server.api_client.kiln_ai_server_client.models import ( GenerateBatchInput, GenerateBatchOutput, - HTTPValidationError, ) -from app.desktop.studio_server.api_client.kiln_ai_server_client.types import Response from app.desktop.studio_server.api_client.kiln_server_client import ( get_authenticated_client, ) @@ -26,6 +23,7 @@ SyntheticDataGenerationSessionConfigApi, TaskInfoApi, ) +from app.desktop.studio_server.utils.response_utils import unwrap_response from fastapi import HTTPException from kiln_ai.datamodel import TaskRun from kiln_ai.datamodel.datamodel_enums import TaskOutputRatingType @@ -93,20 +91,10 @@ async def generate_copilot_examples( body=generate_input, ) ) - check_response_error(detailed_result) - - result = detailed_result.parsed - if result is None: - raise HTTPException( - status_code=500, - detail="Failed to generate synthetic data for spec. Please try again.", - ) - - if isinstance(result, HTTPValidationError): - raise HTTPException( - status_code=422, - detail="Validation error.", - ) + result = unwrap_response( + detailed_result, + none_detail="Failed to generate synthetic data for spec. Please try again.", + ) if not isinstance(result, GenerateBatchOutput): raise HTTPException( @@ -278,23 +266,3 @@ def create_dataset_task_runs( task_runs.append(create_task_run_from_sample(example, train_tag, extra_tags)) return task_runs - - -def check_response_error( - response: Response, default_detail: str = "Unknown error." -) -> None: - """Check if the response is an error with user centric message.""" - if response.status_code != 200: - # response.content is a bytes object - # We check if it's a JSON object with a user message field - detail = default_detail - if response.content.startswith(b"{"): - try: - json_data = json.loads(response.content) - detail = json_data.get("message", default_detail) - except json.JSONDecodeError: - pass - raise HTTPException( - status_code=response.status_code, - detail=detail, - ) diff --git a/app/desktop/studio_server/utils/response_utils.py b/app/desktop/studio_server/utils/response_utils.py new file mode 100644 index 000000000..74a460657 --- /dev/null +++ b/app/desktop/studio_server/utils/response_utils.py @@ -0,0 +1,70 @@ +import json + +from app.desktop.studio_server.api_client.kiln_ai_server_client.models import ( + HTTPValidationError, +) +from app.desktop.studio_server.api_client.kiln_ai_server_client.types import Response +from fastapi import HTTPException +from typing_extensions import TypeVar + + +def check_response_error( + response: Response, default_detail: str = "Unknown error." +) -> None: + """Check if the response is an error with user centric message.""" + if not (200 <= response.status_code < 300): + # response.content is a bytes object + # We check if it's a JSON object with a user message field + detail = default_detail + if response.content.startswith(b"{"): + try: + json_data = json.loads(response.content) + detail = json_data.get("message", default_detail) + except json.JSONDecodeError: + pass + raise HTTPException( + status_code=response.status_code, + detail=detail, + ) + + +T = TypeVar("T") + + +def unwrap_response_allow_none( + response: Response[T | HTTPValidationError], + default_detail: str = "Unknown error.", +) -> T | None: + """ + Raise an error if the response is not 2xx or a validation error, and return the parsed response. + + The returned value is of the type T. + """ + check_response_error(response, default_detail=default_detail) + + parsed_response = response.parsed + # we must check for this to narrow down the type, but this should never + # happen since check_response_error should raise if it is a validation error + if isinstance(parsed_response, HTTPValidationError): + raise RuntimeError("An unknown error occurred.") + + return parsed_response + + +def unwrap_response( + response: Response[T | HTTPValidationError], + default_detail: str = "Unknown error.", + none_detail: str = "An unknown error occurred.", +) -> T: + """ + Raise an error if the response is not 2xx or a validation error or None, and return the parsed response. + + If you want to allow None, use unwrap_response_allow_none instead. + + The returned value is of the type T. + """ + parsed = unwrap_response_allow_none(response, default_detail=default_detail) + if parsed is None: + raise HTTPException(status_code=500, detail=none_detail) + + return parsed diff --git a/app/desktop/studio_server/utils/test_response_utils.py b/app/desktop/studio_server/utils/test_response_utils.py new file mode 100644 index 000000000..d0335dc4b --- /dev/null +++ b/app/desktop/studio_server/utils/test_response_utils.py @@ -0,0 +1,180 @@ +import json +from http import HTTPStatus +from unittest.mock import MagicMock + +import pytest +from app.desktop.studio_server.api_client.kiln_ai_server_client.models import ( + HTTPValidationError, +) +from app.desktop.studio_server.utils.response_utils import ( + check_response_error, + unwrap_response, + unwrap_response_allow_none, +) +from fastapi import HTTPException + + +def _make_response( + status_code: HTTPStatus, content: bytes, parsed: object = None +) -> MagicMock: + resp = MagicMock() + resp.status_code = status_code + resp.content = content + resp.parsed = parsed + return resp + + +@pytest.mark.parametrize( + "status_code", + [HTTPStatus.OK, HTTPStatus.CREATED, HTTPStatus.ACCEPTED, HTTPStatus.NO_CONTENT], +) +def test_2xx_does_not_raise(status_code: HTTPStatus): + resp = _make_response(status_code, b"") + check_response_error(resp) + + +@pytest.mark.parametrize( + "status_code", + [HTTPStatus.BAD_REQUEST, HTTPStatus.INTERNAL_SERVER_ERROR, HTTPStatus.FORBIDDEN], +) +def test_non_200_with_json_message(status_code: HTTPStatus): + body = json.dumps({"message": "Something went wrong"}).encode() + resp = _make_response(status_code, body) + + with pytest.raises(HTTPException) as exc_info: + check_response_error(resp) + + assert exc_info.value.status_code == status_code + assert exc_info.value.detail == "Something went wrong" + + +def test_non_200_with_json_missing_message_uses_default(): + body = json.dumps({"error": "no message key here"}).encode() + resp = _make_response(HTTPStatus.BAD_REQUEST, body) + + with pytest.raises(HTTPException) as exc_info: + check_response_error(resp) + + assert exc_info.value.detail == "Unknown error." + + +def test_non_200_with_custom_default_detail(): + resp = _make_response(HTTPStatus.BAD_REQUEST, b"not json") + + with pytest.raises(HTTPException) as exc_info: + check_response_error(resp, default_detail="Custom default") + + assert exc_info.value.detail == "Custom default" + + +def test_non_200_with_invalid_json_uses_default(): + resp = _make_response(HTTPStatus.BAD_REQUEST, b"{invalid json") + + with pytest.raises(HTTPException) as exc_info: + check_response_error(resp) + + assert exc_info.value.detail == "Unknown error." + + +def test_non_200_with_non_json_content_uses_default(): + resp = _make_response(HTTPStatus.INTERNAL_SERVER_ERROR, b"plain text error") + + with pytest.raises(HTTPException) as exc_info: + check_response_error(resp) + + assert exc_info.value.detail == "Unknown error." + + +# -- unwrap_response_allow_none tests -- + + +def test_unwrap_allow_none_returns_parsed_value(): + resp = _make_response(HTTPStatus.OK, b"", parsed={"key": "value"}) + assert unwrap_response_allow_none(resp) == {"key": "value"} + + +def test_unwrap_allow_none_returns_none_when_parsed_is_none(): + resp = _make_response(HTTPStatus.OK, b"", parsed=None) + assert unwrap_response_allow_none(resp) is None + + +def test_unwrap_allow_none_raises_on_error_status(): + body = json.dumps({"message": "bad request"}).encode() + resp = _make_response(HTTPStatus.BAD_REQUEST, body, parsed=None) + + with pytest.raises(HTTPException) as exc_info: + unwrap_response_allow_none(resp) + + assert exc_info.value.status_code == HTTPStatus.BAD_REQUEST + assert exc_info.value.detail == "bad request" + + +def test_unwrap_allow_none_uses_custom_default_detail(): + resp = _make_response(HTTPStatus.BAD_REQUEST, b"not json", parsed=None) + + with pytest.raises(HTTPException) as exc_info: + unwrap_response_allow_none(resp, default_detail="Custom detail") + + assert exc_info.value.detail == "Custom detail" + + +def test_unwrap_allow_none_raises_on_validation_error(): + resp = _make_response(HTTPStatus.OK, b"", parsed=HTTPValidationError()) + + with pytest.raises(RuntimeError, match="unknown error"): + unwrap_response_allow_none(resp) + + +# -- unwrap_response tests -- + + +def test_unwrap_returns_parsed_value(): + resp = _make_response(HTTPStatus.OK, b"", parsed="hello") + assert unwrap_response(resp) == "hello" + + +def test_unwrap_raises_on_none_parsed(): + resp = _make_response(HTTPStatus.OK, b"", parsed=None) + + with pytest.raises(HTTPException) as exc_info: + unwrap_response(resp) + + assert exc_info.value.status_code == 500 + assert exc_info.value.detail == "An unknown error occurred." + + +def test_unwrap_raises_on_error_status(): + body = json.dumps({"message": "forbidden"}).encode() + resp = _make_response(HTTPStatus.FORBIDDEN, body, parsed=None) + + with pytest.raises(HTTPException) as exc_info: + unwrap_response(resp) + + assert exc_info.value.status_code == HTTPStatus.FORBIDDEN + assert exc_info.value.detail == "forbidden" + + +def test_unwrap_uses_custom_default_detail(): + resp = _make_response(HTTPStatus.BAD_REQUEST, b"not json", parsed=None) + + with pytest.raises(HTTPException) as exc_info: + unwrap_response(resp, default_detail="Custom detail") + + assert exc_info.value.detail == "Custom detail" + + +def test_unwrap_uses_custom_none_detail(): + resp = _make_response(HTTPStatus.OK, b"", parsed=None) + + with pytest.raises(HTTPException) as exc_info: + unwrap_response(resp, none_detail="No data returned") + + assert exc_info.value.status_code == 500 + assert exc_info.value.detail == "No data returned" + + +def test_unwrap_raises_on_validation_error(): + resp = _make_response(HTTPStatus.OK, b"", parsed=HTTPValidationError()) + + with pytest.raises(RuntimeError, match="unknown error"): + unwrap_response(resp)