diff --git a/tools/server/server-common.cpp b/tools/server/server-common.cpp index 8cf794f1e2..a37826abea 100644 --- a/tools/server/server-common.cpp +++ b/tools/server/server-common.cpp @@ -1008,7 +1008,7 @@ server_tokens tokenize_input_subprompt(const llama_vocab * vocab, mtmd_context * return server_tokens(tmp, false); } } else { - throw std::runtime_error("\"prompt\" elements must be a string, a list of tokens, a JSON object containing a prompt string, or a list of mixed strings & tokens."); + throw std::invalid_argument("\"prompt\" elements must be a string, a list of tokens, a JSON object containing a prompt string, or a list of mixed strings & tokens."); } } @@ -1023,7 +1023,7 @@ std::vector tokenize_input_prompts(const llama_vocab * vocab, mtm result.push_back(tokenize_input_subprompt(vocab, mctx, json_prompt, add_special, parse_special, init_opt)); } if (result.empty()) { - throw std::runtime_error("\"prompt\" must not be empty"); + throw std::invalid_argument("\"prompt\" must not be empty"); } return result; } diff --git a/tools/server/server.cpp b/tools/server/server.cpp index 4568a11fb9..bcf84e1ae9 100644 --- a/tools/server/server.cpp +++ b/tools/server/server.cpp @@ -61,6 +61,10 @@ static server_http_context::handler_t ex_wrapper(server_http_context::handler_t // treat invalid_argument as invalid request (400) error = ERROR_TYPE_INVALID_REQUEST; message = e.what(); + } catch (const common_json_error & e) { + // JSON parse and type errors are invalid requests (400) + error = ERROR_TYPE_INVALID_REQUEST; + message = e.what(); } catch (const std::exception & e) { // treat other exceptions as server error (500) error = ERROR_TYPE_SERVER; diff --git a/tools/server/tests/unit/test_embedding.py b/tools/server/tests/unit/test_embedding.py index c4a7d35fef..9efa273fc9 100644 --- a/tools/server/tests/unit/test_embedding.py +++ b/tools/server/tests/unit/test_embedding.py @@ -327,3 +327,21 @@ def test_embedding_openai_library_base64(): # make sure the decoded data is the same as the original for x, y in zip(floats, vec0): assert abs(x - y) < EPSILON + + +@pytest.mark.parametrize( + "data", + [ + {"input": []}, + {"input": True}, + {"input": "hello", "encoding_format": 1}, + ] +) +def test_embedding_invalid_request(data): + global server + server.pooling = 'last' + server.start() + res = server.make_request("POST", "/v1/embeddings", data=data) + assert res.status_code == 400 + assert "error" in res.body + assert res.body["error"]["type"] == "invalid_request_error"