server: return HTTP 400 for invalid embedding requests (#29060)

This commit is contained in:
Sam Malayek
2026-10-01 16:52:43 +02:00
committed by GitHub
parent 2b36825cbc
commit d775ebf363
3 changed files with 24 additions and 2 deletions
+2 -2
View File
@@ -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<server_tokens> 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;
}
+4
View File
@@ -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;
+18
View File
@@ -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"