diff --git a/common/download.cpp b/common/download.cpp index 3776c6c7eb..05d84ca569 100644 --- a/common/download.cpp +++ b/common/download.cpp @@ -309,7 +309,7 @@ static int common_download_file_single_online(const std::string & url, if (!opts.bearer_token.empty()) { headers.emplace("Authorization", "Bearer " + opts.bearer_token); } - cli.set_default_headers(headers); + cli->set_default_headers(headers); std::string last_etag; if (file_exists) { @@ -318,7 +318,7 @@ static int common_download_file_single_online(const std::string & url, LOG_DBG("%s: no previous model file found %s\n", __func__, path.c_str()); } - auto head = cli.Head(parts.path); + auto head = cli->Head(parts.path); if (!head || head->status < 200 || head->status >= 300) { LOG_TRC("%s: HEAD failed, status: %d\n", __func__, head ? head->status : -1); if (file_exists) { @@ -404,7 +404,7 @@ static int common_download_file_single_online(const std::string & url, __func__, common_http_show_masked_url(parts).c_str(), path_temporary.c_str(), etag.c_str()); - if (common_pull_file(cli, parts.path, path_temporary, supports_ranges, p, opts.callback)) { + if (common_pull_file(*cli, parts.path, path_temporary, supports_ranges, p, opts.callback)) { if (std::rename(path_temporary.c_str(), path.c_str()) != 0) { LOG_ERR("%s: unable to rename file: %s to %s\n", __func__, path_temporary.c_str(), path.c_str()); break; @@ -447,12 +447,12 @@ std::pair> common_remote_get_content(const std::string } if (params.timeout > 0) { - cli.set_read_timeout(params.timeout, 0); - cli.set_write_timeout(params.timeout, 0); + cli->set_read_timeout(params.timeout, 0); + cli->set_write_timeout(params.timeout, 0); } std::vector buf; - auto res = cli.Get(parts.path, headers, + auto res = cli->Get(parts.path, headers, [&](const char *data, size_t len) { buf.insert(buf.end(), data, data + len); return params.max_size == 0 || diff --git a/common/hf-cache.cpp b/common/hf-cache.cpp index f1dacaa477..8d6d2bdffb 100644 --- a/common/hf-cache.cpp +++ b/common/hf-cache.cpp @@ -210,7 +210,7 @@ static nl::json api_get(const std::string & url, LOG_WRN("%s: invalid token, authentication disabled\n", __func__); } - if (auto res = cli.Get(parts.path, headers)) { + if (auto res = cli->Get(parts.path, headers)) { auto body = res->body; if (res->status == 200) { diff --git a/common/http.h b/common/http.h index 878ad1ce28..510414412c 100644 --- a/common/http.h +++ b/common/http.h @@ -2,6 +2,8 @@ #include +#include + #ifdef _WIN32 #include #include @@ -97,7 +99,25 @@ static common_http_url common_http_parse_url(const std::string & url) { return parts; } -static std::pair common_http_client(const std::string & url) { +class common_http_client { + httplib::Client cli; +public: + common_http_client(const std::string & url) : cli(url) {} + //set_follow_location // TODO + //set_basic_auth // TODO + httplib::Result Get(); // TODO + // TODO +}; + +using common_http_client_ptr = std::unique_ptr; + +// re-assign this to mock a different HTTP client for testing purposes (ONLY for testing) +inline std::function + common_http_client_factory = [](const std::string & url) -> common_http_client_ptr { + return std::make_unique(url); + }; + +static std::pair common_http_client_init(const std::string & url) { common_http_url parts = common_http_parse_url(url); if (parts.host.empty()) { @@ -115,14 +135,12 @@ static std::pair common_http_client(const std: } #endif - httplib::Client cli(parts.scheme + "://" + common_http_format_host(parts.host) + ":" + std::to_string(parts.port)); + auto cli = common_http_client_factory(parts.scheme + "://" + common_http_format_host(parts.host) + ":" + std::to_string(parts.port)); if (!parts.user.empty()) { - cli.set_basic_auth(parts.user, parts.password); + cli->set_basic_auth(parts.user, parts.password); } - cli.set_follow_location(true); - return { std::move(cli), std::move(parts) }; } diff --git a/tests/gguf-model-data.cpp b/tests/gguf-model-data.cpp index fe8b4ca76e..5a9579dfaf 100644 --- a/tests/gguf-model-data.cpp +++ b/tests/gguf-model-data.cpp @@ -418,13 +418,13 @@ static std::pair> gguf_http_get( auto [cli, parts] = common_http_client(url); if (timeout_sec > 0) { - cli.set_read_timeout(timeout_sec, 0); - cli.set_write_timeout(timeout_sec, 0); + cli->set_read_timeout(timeout_sec, 0); + cli->set_write_timeout(timeout_sec, 0); } - cli.set_connection_timeout(30, 0); + cli->set_connection_timeout(30, 0); std::vector body; - auto res = cli.Get(parts.path, headers, + auto res = cli->Get(parts.path, headers, [&](const char * data, size_t len) { body.insert(body.end(), data, data + len); return true; diff --git a/tools/cli/cli-client.cpp b/tools/cli/cli-client.cpp index 1c563335ba..14aacbdc24 100644 --- a/tools/cli/cli-client.cpp +++ b/tools/cli/cli-client.cpp @@ -27,9 +27,9 @@ static std::string join_path(const common_http_url & parts, const std::string & std::string cli_client::get(const std::string & path) { auto [cli, parts] = common_http_client(server_base); - cli.set_read_timeout(CLI_HTTP_READ_TIMEOUT_SEC, 0); + cli->set_read_timeout(CLI_HTTP_READ_TIMEOUT_SEC, 0); auto path_with_model = path + (model.empty() ? "" : ("?model=" + model)); - auto res = cli.Get(join_path(parts, path_with_model)); + auto res = cli->Get(join_path(parts, path_with_model)); if (!res) { throw std::runtime_error("failed to connect to " + server_base + ": " + httplib::to_string(res.error())); } @@ -41,8 +41,8 @@ std::string cli_client::get(const std::string & path) { std::string cli_client::post(const std::string & path, const std::string & body) { auto [cli, parts] = common_http_client(server_base); - cli.set_read_timeout(CLI_HTTP_READ_TIMEOUT_SEC, 0); - auto res = cli.Post(join_path(parts, path), body, "application/json"); + cli->set_read_timeout(CLI_HTTP_READ_TIMEOUT_SEC, 0); + auto res = cli->Post(join_path(parts, path), body, "application/json"); if (!res) { throw std::runtime_error("failed to connect to " + server_base + ": " + httplib::to_string(res.error())); } @@ -57,7 +57,7 @@ std::string cli_client::post_sse(const std::string & path, const std::function & should_stop, const std::function & on_data) { auto [cli, parts] = common_http_client(server_base); - cli.set_read_timeout(CLI_HTTP_READ_TIMEOUT_SEC, 0); + cli->set_read_timeout(CLI_HTTP_READ_TIMEOUT_SEC, 0); std::string pending; // buffer for incomplete SSE lines std::string raw_body; // accumulated body, used only for error reporting @@ -90,7 +90,7 @@ std::string cli_client::post_sse(const std::string & path, }; httplib::Headers headers = {{"Accept", "text/event-stream"}}; - auto res = cli.Post(join_path(parts, path), headers, body, "application/json", receiver); + auto res = cli->Post(join_path(parts, path), headers, body, "application/json", receiver); if (!res) { if (res.error() == httplib::Error::Canceled && should_stop()) { @@ -111,8 +111,8 @@ bool cli_client::wait_health(const std::function & is_aborted) { int connect_attempts = 0; while (!is_aborted()) { auto [cli, parts] = common_http_client(server_base); - cli.set_connection_timeout(1, 0); - auto res = cli.Get(join_path(parts, "/health")); + cli->set_connection_timeout(1, 0); + auto res = cli->Get(join_path(parts, "/health")); if (res) { if (res->status == 200) { return true; diff --git a/tools/cli/cli-server.h b/tools/cli/cli-server.h index 7596efb01b..bc5d1bc232 100644 --- a/tools/cli/cli-server.h +++ b/tools/cli/cli-server.h @@ -65,8 +65,8 @@ struct cli_server { } while (!should_stop()) { auto [cli, parts] = common_http_client(address()); - cli.set_connection_timeout(1, 0); - auto res = cli.Get("/health"); + cli->set_connection_timeout(1, 0); + auto res = cli->Get("/health"); if (res) { if (res->status == 200) { return true;