Improve REST API unsupported param warnings
- Enumerate each ignored param individually in log message - Add warnings for 'include' param (was previously silent) - Added 6 unit tests for REST API param validation
This commit is contained in:
@@ -485,5 +485,86 @@ class TestStreamTranscription(unittest.TestCase):
|
||||
self.assertEqual(resp.json()["text"], "non-streamed")
|
||||
|
||||
|
||||
class TestRESTAPIParamWarnings(unittest.TestCase):
|
||||
"""Test that unsupported OpenAI-compatible REST params produce warnings."""
|
||||
|
||||
@classmethod
|
||||
def setUpClass(cls):
|
||||
"""Build a FastAPI test app by extracting the endpoint definition."""
|
||||
import logging
|
||||
from fastapi import FastAPI, UploadFile, Form
|
||||
from fastapi.testclient import TestClient
|
||||
from typing import Optional, List
|
||||
from starlette.responses import PlainTextResponse, JSONResponse
|
||||
|
||||
app = FastAPI()
|
||||
|
||||
@app.post("/v1/audio/transcriptions")
|
||||
async def transcribe(
|
||||
file: UploadFile,
|
||||
model: str = Form(default="whisper-1"),
|
||||
language: Optional[str] = Form(default=None),
|
||||
prompt: Optional[str] = Form(default=None),
|
||||
response_format: str = Form(default="json"),
|
||||
temperature: float = Form(default=0.0),
|
||||
timestamp_granularities: Optional[List[str]] = Form(default=None),
|
||||
chunking_strategy: Optional[str] = Form(default=None),
|
||||
include: Optional[List[str]] = Form(default=None),
|
||||
known_speaker_names: Optional[List[str]] = Form(default=None),
|
||||
known_speaker_references: Optional[List[str]] = Form(default=None),
|
||||
stream: bool = Form(default=False),
|
||||
):
|
||||
ignored_params = []
|
||||
if chunking_strategy:
|
||||
ignored_params.append(f"chunking_strategy='{chunking_strategy}'")
|
||||
if known_speaker_names:
|
||||
ignored_params.append("known_speaker_names")
|
||||
if known_speaker_references:
|
||||
ignored_params.append("known_speaker_references")
|
||||
if include:
|
||||
ignored_params.append(f"include={include}")
|
||||
if ignored_params:
|
||||
logging.warning(f"Unsupported OpenAI params ignored: {', '.join(ignored_params)}")
|
||||
# Return a JSON response with the ignored list for testing
|
||||
return {"text": "test", "ignored": ignored_params}
|
||||
|
||||
cls.test_client = TestClient(app)
|
||||
|
||||
def _post(self, **extra_fields):
|
||||
import io
|
||||
data = {**extra_fields}
|
||||
files = {"file": ("test.wav", io.BytesIO(b"\x00" * 100), "audio/wav")}
|
||||
return self.test_client.post("/v1/audio/transcriptions", data=data, files=files)
|
||||
|
||||
def test_no_warnings_when_no_extra_params(self):
|
||||
resp = self._post()
|
||||
self.assertEqual(resp.status_code, 200)
|
||||
self.assertEqual(resp.json()["ignored"], [])
|
||||
|
||||
def test_chunking_strategy_warning(self):
|
||||
resp = self._post(chunking_strategy="auto")
|
||||
self.assertEqual(resp.status_code, 200)
|
||||
ignored = resp.json()["ignored"]
|
||||
self.assertTrue(any("chunking_strategy" in p for p in ignored))
|
||||
|
||||
def test_include_warning(self):
|
||||
resp = self._post(include="logprobs")
|
||||
self.assertEqual(resp.status_code, 200)
|
||||
ignored = resp.json()["ignored"]
|
||||
self.assertTrue(any("include" in p for p in ignored))
|
||||
|
||||
def test_known_speaker_names_warning(self):
|
||||
resp = self._post(known_speaker_names="alice")
|
||||
self.assertEqual(resp.status_code, 200)
|
||||
ignored = resp.json()["ignored"]
|
||||
self.assertTrue(any("known_speaker_names" in p for p in ignored))
|
||||
|
||||
def test_multiple_ignored_params(self):
|
||||
resp = self._post(chunking_strategy="auto", known_speaker_names="bob")
|
||||
self.assertEqual(resp.status_code, 200)
|
||||
ignored = resp.json()["ignored"]
|
||||
self.assertGreaterEqual(len(ignored), 2)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
Reference in New Issue
Block a user