mirror of
https://github.com/ggml-org/llama.cpp.git
synced 2026-10-01 18:37:28 -05:00
server: return HTTP 400 for invalid embedding requests (#29060)
This commit is contained in:
@@ -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;
|
||||
}
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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"
|
||||
|
||||
Reference in New Issue
Block a user