diff --git a/event/hloop.c b/event/hloop.c index 4a0acd814..0c81bcea8 100644 --- a/event/hloop.c +++ b/event/hloop.c @@ -360,10 +360,6 @@ static void hloop_cleanup(hloop_t* loop) { loop->pendings[i] = NULL; } - // async dns resolver - printd("cleanup dns_resolver...\n"); - hdns_resolver_free(loop); - // per-loop lua_State (opaque; destructor supplied by lua/ layer) if (loop->lua_state && loop->lua_state_dtor) { printd("cleanup lua_state...\n"); @@ -390,6 +386,12 @@ static void hloop_cleanup(hloop_t* loop) { } io_array_cleanup(&loop->ios); + // IO close callbacks may cancel pending DNS queries, so keep the resolver + // alive until every IO has been closed. DNS query timers are still valid + // here because the timer sweep happens below. + printd("cleanup dns_resolver...\n"); + hdns_resolver_free(loop); + // idles printd("cleanup idles...\n"); struct list_node* node = loop->idles.next; diff --git a/examples/tcp_proxy_server.c b/examples/tcp_proxy_server.c index 82b60ad8c..57bba3337 100644 --- a/examples/tcp_proxy_server.c +++ b/examples/tcp_proxy_server.c @@ -7,6 +7,7 @@ */ #include "hloop.h" +#include "hsocket.h" int main(int argc, char** argv) { if (argc < 3) { @@ -32,6 +33,17 @@ int main(int argc, char** argv) { setting.target_port = 80; } + if (!is_ipaddr(setting.target_host)) { + sockaddr_u target_addr; + char target_ip[SOCKADDR_STRLEN] = {0}; + if (ResolveAddr(setting.target_host, &target_addr) != 0 || + sockaddr_ip(&target_addr, target_ip, sizeof(target_ip)) == NULL) { + fprintf(stderr, "Could not resolve target host: %s\n", setting.target_host); + return -20; + } + strncpy(setting.target_host, target_ip, sizeof(setting.target_host) - 1); + } + hloop_t* loop = hloop_new(0); hio_t* listener = hloop_create_tcp_proxy_server(loop, &setting); if (listener == NULL) { diff --git a/examples/tinyproxyd.c b/examples/tinyproxyd.c index fe9ba64c0..3a8c8aa58 100644 --- a/examples/tinyproxyd.c +++ b/examples/tinyproxyd.c @@ -41,6 +41,7 @@ static hloop_t** worker_loops = NULL; #define HTTP_KEEPALIVE_TIMEOUT 60000 // ms #define HTTP_MAX_URL_LENGTH 256 #define HTTP_MAX_HEAD_LENGTH 1024 +#define HTTP_CONNECT_RESPONSE "HTTP/1.1 200 Connection Established\r\n\r\n" typedef enum { s_begin, @@ -166,9 +167,16 @@ static bool parse_http_head(http_conn_t* conn, char* buf, int len) { } static void on_upstream_connect(hio_t* upstream_io) { - // printf("on_upstream_connect\n"); http_conn_t* conn = (http_conn_t*)hevent_userdata(upstream_io); http_msg_t* req = &conn->request; + if (stricmp(req->method, "CONNECT") == 0) { + hio_write(conn->io, HTTP_CONNECT_RESPONSE, strlen(HTTP_CONNECT_RESPONSE)); + hio_setcb_read(conn->io, hio_write_upstream); + hio_setcb_read(upstream_io, hio_write_upstream); + hio_read_start(conn->io); + hio_read_start(upstream_io); + return; + } // send head char stackbuf[HTTP_MAX_HEAD_LENGTH + 1024] = {0}; char* buf = stackbuf; @@ -194,17 +202,25 @@ static void on_upstream_connect(hio_t* upstream_io) { static int on_head_end(http_conn_t* conn) { http_msg_t* req = &conn->request; - if (req->host[0] == '\0') { + bool is_connect = stricmp(req->method, "CONNECT") == 0; + const char* authority = is_connect ? req->path : req->host; + if (authority[0] == '\0') { fprintf(stderr, "No Host header!\n"); return -1; } char backend_host[64] = {0}; - strcpy(backend_host, req->host); - int backend_port = 80; - char* pos = strchr(backend_host, ':'); - if (pos) { - *pos = '\0'; - backend_port = atoi(pos + 1); + int backend_port = is_connect ? 443 : 80; + if (authority[0] == '[') { + const char* end = strchr(authority + 1, ']'); + if (end == NULL || end - authority - 1 >= (int)sizeof(backend_host)) return -1; + memcpy(backend_host, authority + 1, end - authority - 1); + if (end[1] == ':') backend_port = atoi(end + 2); + } else { + const char* colon = strrchr(authority, ':'); + size_t host_len = colon ? (size_t)(colon - authority) : strlen(authority); + if (host_len == 0 || host_len >= sizeof(backend_host)) return -1; + memcpy(backend_host, authority, host_len); + if (colon) backend_port = atoi(colon + 1); } if (backend_port == proxy_port && (strcmp(backend_host, proxy_host) == 0 || @@ -215,7 +231,7 @@ static int on_head_end(http_conn_t* conn) { } // NOTE: blew for proxy req->proxy = 1; - int backend_ssl = strncmp(req->path, "https", 5) == 0 ? 1 : 0; + int backend_ssl = !is_connect && strncmp(req->path, "https", 5) == 0 ? 1 : 0; // printf("upstream %s:%d\n", backend_host, backend_port); hloop_t* loop = hevent_loop(conn->io); // hio_t* upstream_io = hio_setup_tcp_upstream(conn->io, backend_host, backend_port, backend_ssl); @@ -327,6 +343,7 @@ static void on_recv(hio_t* io, void* buf, int readbytes) { conn->state = s_end; if (req->proxy) { // NOTE: wait upstream connect! + break; } else { goto s_end; } @@ -361,6 +378,7 @@ static void on_recv(hio_t* io, void* buf, int readbytes) { // received complete request if (req->proxy) { // NOTE: reply by upstream + if (stricmp(req->method, "CONNECT") == 0) break; } else { on_request(conn); } diff --git a/examples/udp_proxy_server.c b/examples/udp_proxy_server.c index c63557eeb..0ca085173 100644 --- a/examples/udp_proxy_server.c +++ b/examples/udp_proxy_server.c @@ -8,6 +8,7 @@ */ #include "hloop.h" +#include "hsocket.h" int main(int argc, char** argv) { if (argc < 3) { @@ -33,6 +34,17 @@ int main(int argc, char** argv) { setting.target_port = 80; } + if (!is_ipaddr(setting.target_host)) { + sockaddr_u target_addr; + char target_ip[SOCKADDR_STRLEN] = {0}; + if (ResolveAddr(setting.target_host, &target_addr) != 0 || + sockaddr_ip(&target_addr, target_ip, sizeof(target_ip)) == NULL) { + fprintf(stderr, "Could not resolve target host: %s\n", setting.target_host); + return -20; + } + strncpy(setting.target_host, target_ip, sizeof(setting.target_host) - 1); + } + hloop_t* loop = hloop_new(0); hio_t* listener = hloop_create_udp_proxy_server(loop, &setting); if (listener == NULL) { diff --git a/http/server/HttpHandler.cpp b/http/server/HttpHandler.cpp index 23d92a671..d93bbb15a 100644 --- a/http/server/HttpHandler.cpp +++ b/http/server/HttpHandler.cpp @@ -34,6 +34,7 @@ HttpHandler::HttpHandler(hio_t* io) : upgrade(0), proxy(0), proxy_connected(0), + proxy_ssl(0), forward_proxy(0), reverse_proxy(0), ip{'\0'}, @@ -53,7 +54,8 @@ HttpHandler::HttpHandler(hio_t* io) : files(NULL), file(NULL), // for proxy - proxy_port(0) + proxy_port(0), + proxy_connect_start_ms(0) { // Init(); } @@ -1167,26 +1169,104 @@ int HttpHandler::connectProxy(const std::string& strUrl) { return SendHttpStatusResponse(HTTP_STATUS_FORBIDDEN); } - hloop_t* loop = hevent_loop(io); proxy = 1; proxy_host = url.host; proxy_port = url.port; - hio_t* upstream_io = hio_create_socket(loop, proxy_host.c_str(), proxy_port, HIO_TYPE_TCP, HIO_CLIENT_SIDE); - if (upstream_io == NULL) { - return SetError(ERR_SOCKET, HTTP_STATUS_BAD_GATEWAY); + proxy_ssl = url.scheme == "https" && req->method != HTTP_CONNECT; + proxy_connect_start_ms = hloop_now_hrtime(hevent_loop(io)) / 1000; + + if (is_ipaddr(proxy_host.c_str())) { + hio_t* upstream_io = hio_create_socket(hevent_loop(io), proxy_host.c_str(), + proxy_port, HIO_TYPE_TCP, HIO_CLIENT_SIDE); + if (upstream_io == NULL) { + return SetError(ERR_SOCKET, HTTP_STATUS_BAD_GATEWAY); + } + return connectProxy(upstream_io); + } + hio_read_stop(io); + + EventLoop* loop = currentThreadEventLoop; + assert(loop != NULL); + + hdns_setting_t dns_setting; + if (service->proxy_connect_timeout > 0) { + int dns_attempts = HDNS_DEFAULT_RETRIES + 1; + dns_setting.timeout_ms = MAX(1, service->proxy_connect_timeout / dns_attempts); + } + hio_t* downstream_io = io; + uint32_t downstream_id = hio_id(io); + std::string target_host = proxy_host; + int target_port = proxy_port; + DnsID dns_id = loop->resolveDns(target_host.c_str(), + [downstream_io, downstream_id, target_host, target_port](int status, int naddrs, const sockaddr_u* addrs) { + if (!hio_is_opened(downstream_io) || hio_id(downstream_io) != downstream_id) { + return; + } + HttpHandler* handler = (HttpHandler*)hevent_userdata(downstream_io); + if (handler == NULL || handler->proxy_host != target_host || + handler->proxy_port != target_port) return; + if (status != HDNS_STATUS_OK || naddrs <= 0) { + int dns_error = status == HDNS_STATUS_TIMEOUT ? ETIMEDOUT : ERR_DNS_RESOLVE; + http_status http_error = dns_error == ETIMEDOUT ? + HTTP_STATUS_GATEWAY_TIMEOUT : HTTP_STATUS_BAD_GATEWAY; + handler->SetError(dns_error, http_error); + handler->SendHttpStatusResponse(http_error); + hio_close(downstream_io); + return; + } + + hio_t* upstream_io = NULL; + char resolved_ip[SOCKADDR_STRLEN] = {0}; + for (int i = 0; i < naddrs; ++i) { + if (sockaddr_ip((sockaddr_u*)&addrs[i], resolved_ip, sizeof(resolved_ip)) == NULL) continue; + upstream_io = hio_create_socket(hevent_loop(downstream_io), resolved_ip, + target_port, HIO_TYPE_TCP, HIO_CLIENT_SIDE); + if (upstream_io) break; + } + if (upstream_io == NULL) { + handler->SetError(ERR_SOCKET, HTTP_STATUS_BAD_GATEWAY); + handler->SendHttpStatusResponse(HTTP_STATUS_BAD_GATEWAY); + hio_close(downstream_io); + return; + } + hio_set_hostname(upstream_io, target_host.c_str()); + handler->connectProxy(upstream_io); + }, &dns_setting); + if (dns_id == INVALID_DNS_ID) { + SetError(ERR_DNS_RESOLVE, HTTP_STATUS_BAD_GATEWAY); + SendHttpStatusResponse(HTTP_STATUS_BAD_GATEWAY); + hio_close(io); + return ERR_DNS_RESOLVE; + } + return 0; +} + +int HttpHandler::connectProxy(hio_t* upstream_io) { + if (!io || !upstream_io) return ERR_NULL_POINTER; + int timeout_ms = service->proxy_connect_timeout; + if (timeout_ms > 0) { + uint64_t elapsed_ms = hloop_now_hrtime(hevent_loop(io)) / 1000 - proxy_connect_start_ms; + if (elapsed_ms >= (uint64_t)timeout_ms) { + hio_close(upstream_io); + SetError(ETIMEDOUT, HTTP_STATUS_GATEWAY_TIMEOUT); + SendHttpStatusResponse(HTTP_STATUS_GATEWAY_TIMEOUT); + hio_close(io); + return ETIMEDOUT; + } + timeout_ms -= (int)elapsed_ms; } // CONNECT establishes a raw TCP tunnel. The client starts TLS only after // receiving the 200 response, so enabling TLS on this upstream would // incorrectly terminate and re-encrypt the tunnel. - if (url.scheme == "https" && req->method != HTTP_CONNECT) { + if (proxy_ssl) { hio_enable_ssl(upstream_io); } hevent_set_userdata(upstream_io, this); hio_setup_upstream(io, upstream_io); hio_setcb_connect(upstream_io, HttpHandler::onProxyConnect); hio_setcb_close(upstream_io, HttpHandler::onProxyClose); - if (service->proxy_connect_timeout > 0) { - hio_set_connect_timeout(upstream_io, service->proxy_connect_timeout); + if (timeout_ms > 0) { + hio_set_connect_timeout(upstream_io, timeout_ms); } if (service->proxy_read_timeout > 0) { hio_set_read_timeout(upstream_io, service->proxy_read_timeout); @@ -1201,7 +1281,7 @@ int HttpHandler::connectProxy(const std::string& strUrl) { } int HttpHandler::closeProxy() { - if (proxy && proxy_connected) { + if (proxy) { proxy_connected = 0; if (io) hio_close_upstream(io); } diff --git a/http/server/HttpHandler.h b/http/server/HttpHandler.h index aaa45b34b..7109cda9a 100644 --- a/http/server/HttpHandler.h +++ b/http/server/HttpHandler.h @@ -38,6 +38,7 @@ class HttpHandler { unsigned upgrade :1; unsigned proxy :1; unsigned proxy_connected :1; + unsigned proxy_ssl :1; unsigned forward_proxy :1; unsigned reverse_proxy :1; @@ -82,6 +83,7 @@ class HttpHandler { // for proxy std::string proxy_host; int proxy_port; + uint64_t proxy_connect_start_ms; HttpHandler(hio_t* io = NULL); ~HttpHandler(); @@ -178,6 +180,7 @@ class HttpHandler { int handleForwardProxy(); int handleReverseProxy(); int connectProxy(const std::string& url); + int connectProxy(hio_t* upstream_io); int closeProxy(); int sendProxyRequest(); static void onProxyConnect(hio_t* upstream_io); diff --git a/unittest/hdns_test.c b/unittest/hdns_test.c index 97783ff51..0dd44623e 100644 --- a/unittest/hdns_test.c +++ b/unittest/hdns_test.c @@ -14,6 +14,7 @@ * 6. NXDOMAIN handling * 7. cancel before completion (callback not invoked) * 8. auto nameserver list is refreshed (throttled) between resolves + * 9. IO close callback can cancel pending DNS during loop cleanup */ #include @@ -174,6 +175,42 @@ static void stop_after(htimer_t* timer) { hloop_stop(hevent_loop(timer)); } +typedef struct { + hdns_t* query; +} cleanup_cancel_ctx_t; + +static void cleanup_cancel_close(hio_t* io) { + cleanup_cancel_ctx_t* ctx = (cleanup_cancel_ctx_t*)hevent_userdata(io); + if (ctx->query) { + hdns_cancel(ctx->query); + ctx->query = NULL; + } +} + +static void test_cleanup_cancel_pending_dns(void) { + hloop_t* loop = hloop_new(0); + cleanup_cancel_ctx_t ctx; + memset(&ctx, 0, sizeof(ctx)); + + hio_t* io = hloop_create_udp_server(loop, "127.0.0.1", 0); + assert(io != NULL); + hevent_set_userdata(io, &ctx); + hio_setcb_close(io, cleanup_cancel_close); + + hdns_setting_t opt; + memset(&opt, 0, sizeof(opt)); + opt.family = HDNS_QUERY_A; + opt.timeout_ms = 10000; + opt.retries = 0; + opt.use_cache = 0; + opt.nameserver = "127.0.0.1:1"; + ctx.query = hdns_resolve_ex(loop, "cleanup.test", &opt, on_never, NULL); + assert(ctx.query != NULL); + + hloop_free(&loop); + assert(ctx.query == NULL); +} + int main() { hloop_t* loop = hloop_new(0); @@ -270,6 +307,8 @@ int main() { hio_close(mock); hloop_free(&loop); + + test_cleanup_cancel_pending_dns(); printf("\nALL hdns_test PASSED\n"); return 0; }