Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 2 additions & 2 deletions tools/server/server-common.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -1008,7 +1008,7 @@ static server_tokens tokenize_input_subprompt(const llama_vocab * vocab, mtmd_co
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.");
}
}

Expand All @@ -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;
}
Expand Down
4 changes: 4 additions & 0 deletions tools/server/server.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down
18 changes: 18 additions & 0 deletions tools/server/tests/unit/test_embedding.py
Original file line number Diff line number Diff line change
Expand Up @@ -289,3 +289,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"
Loading