wormhole-proxy 3.2.2__tar.gz → 3.3.0__tar.gz

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.3
2
2
  Name: wormhole-proxy
3
- Version: 3.2.2
3
+ Version: 3.3.0
4
4
  Summary: Asynchronous I/O HTTP and HTTPS Proxy on Python >= 3.11
5
5
  License: MIT
6
6
  Keywords: wormhole,asynchronous,web,proxy
@@ -19,6 +19,7 @@ Classifier: Programming Language :: Python :: 3.12
19
19
  Classifier: Programming Language :: Python :: 3.13
20
20
  Classifier: Topic :: Internet :: Proxy Servers
21
21
  Provides-Extra: performance
22
+ Requires-Dist: aiodns (>=3.5.0,<4.0.0)
22
23
  Requires-Dist: aiohttp (>=3.12.13,<4.0.0)
23
24
  Requires-Dist: aiosqlite (>=0.17.0,<0.18.0)
24
25
  Requires-Dist: loguru (>=0.7.2,<0.8.0)
@@ -59,6 +60,7 @@ Description-Content-Type: text/markdown
59
60
  ## Dependency
60
61
 
61
62
  - Python \>= 3.11
63
+ - aiodns
62
64
  - aiohttp
63
65
  - aiosqlite
64
66
  - loguru
@@ -29,6 +29,7 @@
29
29
  ## Dependency
30
30
 
31
31
  - Python \>= 3.11
32
+ - aiodns
32
33
  - aiohttp
33
34
  - aiosqlite
34
35
  - loguru
@@ -1,7 +1,7 @@
1
1
  # Main project metadata (PEP 621 standard)
2
2
  [project]
3
3
  name = "wormhole-proxy"
4
- version = "3.2.2"
4
+ version = "3.3.0"
5
5
  description = "Asynchronous I/O HTTP and HTTPS Proxy on Python >= 3.11"
6
6
  readme = "README.md"
7
7
  authors = [
@@ -43,6 +43,7 @@ performance = ["winloop", "uvloop"]
43
43
  packages = [{include = "wormhole"}]
44
44
 
45
45
  [tool.poetry.dependencies]
46
+ aiodns = "^3.5.0"
46
47
  aiohttp = "^3.12.13"
47
48
  aiosqlite = "^0.17.0"
48
49
  loguru = "^0.7.2"
@@ -7,7 +7,7 @@ import asyncio
7
7
  import re
8
8
 
9
9
  # A curated list of popular and well-maintained ad-block lists
10
- BLOCKLIST_URLS = [
10
+ BLOCKLIST_URLS: list[str] = [
11
11
  "https://raw.githubusercontent.com/StevenBlack/hosts/master/hosts",
12
12
  "https://pgl.yoyo.org/adservers/serverlist.php?hostformat=hosts&showintro=0&mimetype=plaintext",
13
13
  "https://easylist.to/easylist/easylist.txt",
@@ -25,7 +25,7 @@ DOMAIN_REGEX = re.compile(
25
25
 
26
26
  async def _fetch_list(session: aiohttp.ClientSession, url: str) -> str:
27
27
  """Fetches the content of a single blocklist URL with a retry mechanism."""
28
- max_retries = 3
28
+ max_retries: int = 3
29
29
  timeout = aiohttp.ClientTimeout(total=15) # seconds
30
30
  for attempt in range(max_retries):
31
31
  try:
@@ -79,8 +79,8 @@ def _filter_redundant_domains(domains: set[str]) -> set[str]:
79
79
  e.g., if 'example.com' is present, 'ad.example.com' is removed.
80
80
  """
81
81
  # Sort by length descending to ensure we process subdomains before parents
82
- sorted_domains = sorted(list(domains), key=len, reverse=True)
83
- optimized_set = set(sorted_domains)
82
+ sorted_domains: list[str] = sorted(list(domains), key=len, reverse=True)
83
+ optimized_set: set[str] = set(sorted_domains)
84
84
 
85
85
  for domain in sorted_domains:
86
86
  parts = domain.split(".")
@@ -106,7 +106,7 @@ async def update_database(
106
106
  db_path.parent.mkdir(parents=True, exist_ok=True)
107
107
 
108
108
  # Start with the hardcoded default allowlist
109
- allowlist_domains = DEFAULT_ALLOWLIST.copy()
109
+ allowlist_domains: set[str] = DEFAULT_ALLOWLIST.copy()
110
110
  logger.info(
111
111
  f"Loaded {len(allowlist_domains)} domains from the default allowlist."
112
112
  )
@@ -152,8 +152,10 @@ async def update_database(
152
152
 
153
153
  # Optimize the final blocklist
154
154
  logger.info("Optimizing list by removing redundant subdomains...")
155
- optimized_domains = _filter_redundant_domains(all_blocked_domains)
156
- num_redundant_removed = len(all_blocked_domains) - len(optimized_domains)
155
+ optimized_domains: set[str] = _filter_redundant_domains(all_blocked_domains)
156
+ num_redundant_removed: int = len(all_blocked_domains) - len(
157
+ optimized_domains
158
+ )
157
159
  logger.info(
158
160
  f"Optimization complete. Removed {num_redundant_removed} redundant domains."
159
161
  )
@@ -1,19 +1,20 @@
1
1
  from pathlib import Path
2
+ from typing import Any
2
3
  import getpass
3
4
  import hashlib
4
5
  import os
5
6
  import stat
6
7
  import sys
7
8
 
8
- REALM = "Wormhole Proxy"
9
- HASH_ALGORITHM = hashlib.sha256 # Use SHA-256
9
+ REALM: str = "Wormhole Proxy"
10
+ HASH_ALGORITHM = hashlib.sha256
10
11
 
11
12
 
12
13
  def _get_password_confirm() -> str | None:
13
14
  """Gets and confirms a new password from the user."""
14
15
  try:
15
- p1 = getpass.getpass()
16
- p2 = getpass.getpass("Retype password: ")
16
+ p1: str = getpass.getpass()
17
+ p2: str = getpass.getpass("Retype password: ")
17
18
  if p1 != p2:
18
19
  print("Passwords do not match.", file=sys.stderr)
19
20
  return None
@@ -61,9 +62,9 @@ def _secure_create_file(path: Path) -> bool:
61
62
  return False
62
63
 
63
64
 
64
- def _read_auth_file(path: Path) -> dict:
65
+ def _read_auth_file(path: Path) -> dict[str, Any]:
65
66
  """Reads the auth file into a dictionary."""
66
- users = {}
67
+ users: dict[str, Any] = {}
67
68
  if path.is_file():
68
69
  with open(path, "r", encoding="utf-8") as f:
69
70
  for line in f:
@@ -5,7 +5,7 @@ import re
5
5
  import secrets
6
6
 
7
7
  # This must match the REALM in auth_manager.py
8
- REALM = "Wormhole Proxy"
8
+ REALM: str = "Wormhole Proxy"
9
9
  HASH_ALGORITHM = hashlib.sha256
10
10
 
11
11
  # Caches for performance
@@ -46,10 +46,13 @@ def _load_auth_file(path: Path) -> dict | None:
46
46
  return _auth_file_cache
47
47
 
48
48
 
49
+ # This regex handles quoted and unquoted values
50
+ QUOTE_UNQUOTE_RE = re.compile(r'(\w+)=(?:"([^"]*)"|([^\s,]*))')
51
+
52
+
49
53
  def _parse_digest_header(header_value: str) -> dict[str, str]:
50
54
  """Parses the Digest authentication header into a dictionary."""
51
- # This regex handles quoted and unquoted values
52
- parts = re.findall(r'(\w+)=(?:"([^"]*)"|([^\s,]*))', header_value)
55
+ parts = QUOTE_UNQUOTE_RE.findall(header_value)
53
56
  # The regex produces tuples like ('key', 'quoted_val', ''), so we merge
54
57
  return {key: val1 or val2 for key, val1, val2 in parts}
55
58
 
@@ -80,7 +83,6 @@ async def verify_credentials(
80
83
  reader: asyncio.StreamReader,
81
84
  writer: asyncio.StreamWriter,
82
85
  method: str, # HTTP Method (e.g., 'CONNECT') is needed for HA2 calculation
83
- uri: str, # The original URI from the request line
84
86
  headers: list[str],
85
87
  auth_file_path: str,
86
88
  ) -> dict[str, str] | None:
@@ -1,6 +1,7 @@
1
1
  from .logger import logger, format_log_message as flm
2
2
  from .safeguards import has_public_ipv6, is_ad_domain, is_private_ip
3
3
  from .tools import get_content_length, get_host_and_port
4
+ from .resolver import resolver
4
5
  import asyncio
5
6
  import ipaddress
6
7
  import random
@@ -66,6 +67,7 @@ async def _resolve_and_validate_host(
66
67
  ) -> list[str]:
67
68
  """
68
69
  Resolves a hostname to a list of IPs, validates them, and caches the list.
70
+ Uses aiodns via a custom resolver that respects the local hosts file.
69
71
  Supports DNS load balancing and prioritizes IPv6 if available.
70
72
  Raises:
71
73
  PermissionError: If the host is an ad domain or resolves to only private IPs.
@@ -88,11 +90,9 @@ async def _resolve_and_validate_host(
88
90
  )
89
91
  return ip_list
90
92
 
91
- # Resolve hostname
92
- loop = asyncio.get_running_loop()
93
+ # Resolve hostname using aiodns resolver
93
94
  try:
94
- addr_info_list = await loop.getaddrinfo(host, None, family=0)
95
- resolved_ips = {info[4][0] for info in addr_info_list}
95
+ resolved_ips = await resolver.resolve(host)
96
96
  except OSError as e:
97
97
  raise OSError(f"Failed to resolve host: {host}") from e
98
98
 
@@ -153,58 +153,93 @@ async def _resolve_and_validate_host(
153
153
  return final_ip_list
154
154
 
155
155
 
156
- async def _create_connection_with_retries(
156
+ async def _create_fastest_connection(
157
157
  ip_list: list[str],
158
158
  port: int,
159
159
  ident: dict[str, str],
160
- max_attempts: int = 3,
161
160
  timeout: int = 5,
161
+ max_attempts: int = 3,
162
162
  verbose: int = 0,
163
163
  ) -> tuple[asyncio.StreamReader, asyncio.StreamWriter]:
164
164
  """
165
- Tries to connect to a list of IPs with a fast timeout and retry mechanism.
165
+ Implements a robust "Happy Eyeballs" connection algorithm with retries.
166
+ It wraps the concurrent connection logic in a retry loop.
166
167
  """
167
168
  last_error = None
168
169
 
169
- # Create a list of connection targets to try, ensuring we don't exceed max_attempts
170
- targets_to_try = (ip_list * (max_attempts // len(ip_list) + 1))[
171
- :max_attempts
172
- ]
173
-
174
- for i, ip in enumerate(targets_to_try):
175
- attempt = i + 1
170
+ for attempt in range(max_attempts):
176
171
  logger.debug(
177
172
  flm(
178
- f"Connection attempt {attempt}/{max_attempts} to {ip}:{port}",
173
+ f"Connection attempt {attempt + 1}/{max_attempts} to {ip_list}",
179
174
  ident,
180
175
  verbose,
181
176
  )
182
177
  )
183
- try:
184
- # Use a short timeout for each connection attempt
185
- reader, writer = await asyncio.wait_for(
186
- asyncio.open_connection(ip, port), timeout=timeout
178
+
179
+ tasks = {
180
+ asyncio.create_task(
181
+ asyncio.wait_for(
182
+ asyncio.open_connection(ip, port), timeout=timeout
183
+ ),
184
+ name=ip,
187
185
  )
188
- logger.debug(
189
- flm(
190
- f"Successfully connected to {ip}:{port} on attempt {attempt}",
191
- ident,
192
- verbose,
193
- )
186
+ for ip in ip_list
187
+ }
188
+
189
+ # Inner loop for the "Happy Eyeballs" race
190
+ while tasks:
191
+ done, pending = await asyncio.wait(
192
+ tasks, return_when=asyncio.FIRST_COMPLETED
194
193
  )
195
- return reader, writer
196
- except (OSError, asyncio.TimeoutError) as e:
197
- last_error = e
194
+
195
+ for task in done:
196
+ try:
197
+ reader, writer = task.result()
198
+ # On success, cancel pending tasks and return the connection
199
+ for p_task in pending:
200
+ p_task.cancel()
201
+ if pending:
202
+ await asyncio.gather(*pending, return_exceptions=True)
203
+ peer = writer.get_extra_info("peername")
204
+ logger.debug(
205
+ flm(
206
+ f"Successfully established fastest connection to {peer[0]}:{peer[1]}",
207
+ ident,
208
+ verbose,
209
+ )
210
+ )
211
+ return reader, writer
212
+ except (
213
+ OSError,
214
+ asyncio.TimeoutError,
215
+ asyncio.CancelledError,
216
+ ) as e:
217
+ ip = task.get_name()
218
+ logger.debug(
219
+ flm(
220
+ f"Connection to {ip}:{port} failed within race: {e}",
221
+ ident,
222
+ verbose,
223
+ )
224
+ )
225
+ last_error = e
226
+
227
+ tasks = pending
228
+
229
+ # If the inner loop finishes, all IPs failed in this attempt.
230
+ # Wait before the next retry, if any.
231
+ if attempt < max_attempts - 1:
198
232
  logger.warning(
199
233
  flm(
200
- f"Connection to {ip}:{port} failed on attempt {attempt}: {e}",
234
+ f"All connections failed on attempt {attempt + 1}. Retrying in 1 second...",
201
235
  ident,
202
236
  verbose,
203
237
  )
204
238
  )
239
+ await asyncio.sleep(1)
205
240
 
206
241
  raise OSError(
207
- f"Failed to connect after {max_attempts} attempts. Last error: {last_error}"
242
+ f"All connection attempts failed after {max_attempts} retries. Last error: {last_error}"
208
243
  )
209
244
 
210
245
 
@@ -218,6 +253,7 @@ async def process_https_tunnel(
218
253
  uri: str,
219
254
  ident: dict[str, str],
220
255
  allow_private: bool,
256
+ max_attempts: int = 3,
221
257
  verbose: int = 0,
222
258
  ) -> None:
223
259
  """Establishes an HTTPS tunnel and relays data between client and server."""
@@ -230,10 +266,8 @@ async def process_https_tunnel(
230
266
  ip_list = await _resolve_and_validate_host(
231
267
  host, ident, allow_private, verbose
232
268
  )
233
-
234
- # Attempt to connect to one of the IPs with retry logic.
235
- server_reader, server_writer = await _create_connection_with_retries(
236
- ip_list, port, ident, verbose=verbose
269
+ server_reader, server_writer = await _create_fastest_connection(
270
+ ip_list, port, ident, max_attempts=max_attempts, verbose=verbose
237
271
  )
238
272
 
239
273
  # Signal the client that the tunnel is established.
@@ -259,6 +293,7 @@ async def process_https_tunnel(
259
293
  logger.warning(flm(f"{method} 403 {uri} ({e})", ident, verbose))
260
294
  client_writer.write(b"HTTP/1.1 403 Forbidden\r\n\r\n")
261
295
  await client_writer.drain()
296
+
262
297
  except Exception as e:
263
298
  msg = flm(f"{method} 502 {uri} ({e})", ident, verbose)
264
299
  if verbose > 2: # Show full traceback only for -vv
@@ -288,10 +323,8 @@ async def _send_http_request(
288
323
  """Helper function to connect and send an HTTP request."""
289
324
  request_line = f"{method} {path or '/'} {version}".encode()
290
325
  headers_bytes = "\r\n".join(headers).encode()
291
-
292
- # Attempt to connect to one of the IPs with retry logic.
293
- server_reader, server_writer = await _create_connection_with_retries(
294
- ip_list, port, ident, max_attempts, verbose=verbose
326
+ server_reader, server_writer = await _create_fastest_connection(
327
+ ip_list, port, ident, max_attempts=max_attempts, verbose=verbose
295
328
  )
296
329
 
297
330
  server_writer.write(request_line + b"\r\n" + headers_bytes + b"\r\n\r\n")
@@ -429,13 +462,16 @@ async def process_http_request(
429
462
  final_headers = [
430
463
  h for h in headers if not h.lower().startswith("proxy-")
431
464
  ]
465
+
432
466
  if not any(h.lower().startswith("host:") for h in final_headers):
433
467
  final_headers.insert(0, f"Host: {host_header}")
468
+
434
469
  final_headers = [
435
470
  h
436
471
  for h in final_headers
437
472
  if not h.lower().startswith("connection:")
438
473
  ]
474
+
439
475
  final_headers.append("Connection: close")
440
476
 
441
477
  server_reader, server_writer = await _send_http_request(
@@ -472,6 +508,7 @@ async def process_http_request(
472
508
  logger.warning(flm(f"{method} 403 {uri} ({e})", ident, verbose))
473
509
  client_writer.write(b"HTTP/1.1 403 Forbidden\r\n\r\n")
474
510
  await client_writer.drain()
511
+
475
512
  except Exception as e:
476
513
  msg = flm(f"{method} 502 {uri} ({e})", ident, verbose)
477
514
  if verbose > 2: # Show full traceback only for -vv
@@ -484,6 +521,7 @@ async def process_http_request(
484
521
  await client_writer.drain()
485
522
  except ConnectionError:
486
523
  pass # Ignore if client is already closed
524
+
487
525
  finally:
488
526
  if server_writer and not server_writer.is_closing():
489
527
  server_writer.close()
@@ -22,7 +22,8 @@ class LogThrottler:
22
22
  if self.repeat_count > 2:
23
23
  self.logger.opt(depth=2).log(
24
24
  self.level,
25
- f"{self.last_message} (and {self.repeat_count -1} more in the last {self.delay} seconds.)",
25
+ f"{self.last_message} "
26
+ f"(and {self.repeat_count -1} more in the last {self.delay} seconds.)",
26
27
  **kwargs,
27
28
  )
28
29
  elif self.repeat_count == 2:
@@ -62,7 +63,10 @@ class LogThrottler:
62
63
  # In loguru, the logger is imported and ready to be configured.
63
64
  # We just need to ensure other modules import this configured instance.
64
65
  def setup_logger(
65
- syslog_host: str | None = None, syslog_port: int = 514, verbose: int = 0
66
+ syslog_host: str | None = None,
67
+ syslog_port: int = 514,
68
+ verbose: int = 0,
69
+ async_mode: bool = True,
66
70
  ) -> None:
67
71
  """
68
72
  Configures the global loguru logger instance. This should only be called once.
@@ -110,10 +114,12 @@ def setup_logger(
110
114
  logging.getLogger("asyncio").setLevel(
111
115
  logging.DEBUG if verbose >= 2 else logging.CRITICAL
112
116
  )
113
- if verbose < 2:
114
- logger.info = LogThrottler(logger, "info").process
115
- logger.warning = LogThrottler(logger, "warning").process
116
- logger.error = LogThrottler(logger, "error").process
117
+
118
+ # Only enable the async LogThrottler if we are in async mode.
119
+ if async_mode and verbose < 2:
120
+ logger.info = LogThrottler(logger, "info").process # type: ignore
121
+ logger.warning = LogThrottler(logger, "warning").process # type: ignore
122
+ logger.error = LogThrottler(logger, "error").process # type: ignore
117
123
 
118
124
 
119
125
  def format_log_message(
@@ -10,6 +10,7 @@ if sys.version_info < (3, 11):
10
10
  from .ad_blocker import update_database
11
11
  from .auth_manager import add_user, modify_user, delete_user
12
12
  from .logger import logger, setup_logger, format_log_message as flm
13
+ from .resolver import resolver
13
14
  from .safeguards import load_ad_block_db, load_allowlist
14
15
  from .server import start_wormhole_server
15
16
  from .version import VERSION
@@ -210,7 +211,8 @@ def main() -> int:
210
211
  )
211
212
  args = parser.parse_args()
212
213
 
213
- # Print license information and exit.
214
+ # --- Utility Command Handling ---
215
+ # These commands run synchronously and exit.
214
216
  if args.license:
215
217
  print(parser.description)
216
218
  try:
@@ -222,7 +224,6 @@ def main() -> int:
222
224
  return 1
223
225
  return 0
224
226
 
225
- # Handle authentication management commands.
226
227
  if args.auth_add:
227
228
  return add_user(args.auth_add[0], args.auth_add[1])
228
229
  if args.auth_mod:
@@ -230,12 +231,22 @@ def main() -> int:
230
231
  if args.auth_del:
231
232
  return delete_user(args.auth_del[0], args.auth_del[1])
232
233
 
233
- # Setup the uvloop before any other operations thay migh use the event loop.
234
+ # --- Server and DB Update Handling ---
235
+ # Determine if we are running the full async server or a utility.
236
+ is_server_mode = not args.update_ad_block_db
237
+
238
+ # Setup the uvloop before any other operations that might use the event loop.
234
239
  if uvloop:
235
240
  uvloop.install()
236
241
 
237
- # Setup logging before any other operations.
238
- setup_logger(args.syslog_host, args.syslog_port, args.verbose)
242
+ # Setup logging. Disable async features for synchronous utility commands.
243
+ setup_logger(
244
+ args.syslog_host,
245
+ args.syslog_port,
246
+ args.verbose,
247
+ async_mode=is_server_mode,
248
+ )
249
+ resolver.configure(verbose=args.verbose)
239
250
 
240
251
  if args.update_ad_block_db:
241
252
  # For this standalone utility, configure a simple logger to show progress.
@@ -249,6 +260,7 @@ def main() -> int:
249
260
  return 1
250
261
  return 0
251
262
 
263
+ # --- Main Server Execution ---
252
264
  if not 1024 <= args.port <= 65535:
253
265
  parser.error("Port must be between 1024 and 65535.")
254
266
 
@@ -0,0 +1,173 @@
1
+ from .logger import logger, format_log_message as flm
2
+ from pathlib import Path
3
+ from typing import ClassVar
4
+ import aiodns
5
+ import asyncio
6
+ import os
7
+ import re
8
+ import sys
9
+
10
+ # Regex to parse a line in the hosts file.
11
+ # It captures the IP address and all hostnames on the line.
12
+ # It ignores comments starting with '#'.
13
+ HOSTS_LINE_REGEX = re.compile(r"^\s*([^\s#]+)\s+([^#]+)")
14
+
15
+
16
+ class Resolver:
17
+ """
18
+ A DNS resolver that uses aiodns and a persistent hosts file cache.
19
+ This class is implemented as a singleton to ensure only one instance
20
+ is used throughout the application.
21
+ """
22
+
23
+ # Class variable to hold the singleton instance
24
+ _instance: ClassVar["Resolver | None"] = None
25
+
26
+ def __init__(self) -> None:
27
+ """
28
+ Initializes the resolver, loads the hosts file, and sets the
29
+ singleton instance. The aiodns.DNSResolver is created lazily.
30
+ """
31
+ if not Resolver._instance:
32
+ self.resolver: aiodns.DNSResolver | None = None
33
+ self.hosts_cache: dict[str, str] = {}
34
+ self.verbose: int = 0 # Verbosity level, configured separately
35
+ self._load_hosts_file()
36
+ Resolver._instance = self
37
+
38
+ def configure(self, verbose: int) -> None:
39
+ """Sets the verbosity level for logging."""
40
+ self.verbose = verbose
41
+
42
+ @staticmethod
43
+ def get_instance() -> "Resolver":
44
+ """Static access method to get the singleton instance."""
45
+ if Resolver._instance is None:
46
+ Resolver()
47
+ return Resolver._instance # type: ignore
48
+
49
+ def _get_hosts_path(self) -> Path:
50
+ """Determines the correct path to the hosts file based on the OS."""
51
+ if sys.platform == "win32":
52
+ # Use the SYSTEMROOT environment variable on Windows
53
+ return (
54
+ Path(os.environ["SYSTEMROOT"])
55
+ / "System32"
56
+ / "drivers"
57
+ / "etc"
58
+ / "hosts"
59
+ )
60
+ else:
61
+ # For Linux, macOS, and other UNIX-like systems
62
+ return Path("/etc/hosts")
63
+
64
+ def _load_hosts_file(self) -> None:
65
+ """Parses the system's hosts file and populates the cache."""
66
+ hosts_path = self._get_hosts_path()
67
+ ident = {"id": "resolver", "client": "internal"}
68
+
69
+ if not hosts_path.exists():
70
+ logger.warning(
71
+ flm(
72
+ f"Hosts file not found at {hosts_path}, skipping.",
73
+ ident,
74
+ self.verbose,
75
+ )
76
+ )
77
+ return
78
+
79
+ try:
80
+ with open(hosts_path, "r", encoding="utf-8") as f:
81
+ for line in f:
82
+ match = HOSTS_LINE_REGEX.match(line)
83
+ if not match:
84
+ continue
85
+
86
+ ip_address = match.group(1)
87
+ hostnames_str = match.group(2)
88
+ hostnames = hostnames_str.strip().split()
89
+
90
+ for hostname in hostnames:
91
+ self.hosts_cache[hostname.lower()] = ip_address
92
+
93
+ logger.info(
94
+ flm(
95
+ f"Loaded {len(self.hosts_cache)} entries from {hosts_path}",
96
+ ident,
97
+ self.verbose,
98
+ )
99
+ )
100
+ except Exception as e:
101
+ logger.error(
102
+ flm(
103
+ f"Failed to load or parse hosts file at {hosts_path}: {e}",
104
+ ident,
105
+ self.verbose,
106
+ )
107
+ )
108
+
109
+ async def resolve(self, hostname: str) -> list[str]:
110
+ """
111
+ Resolves a hostname to a list of IP addresses.
112
+
113
+ 1. Checks the local hosts file cache first.
114
+ 2. If not found, queries DNS for A and AAAA records using aiodns.
115
+ """
116
+ # Create a specific ident for this resolution request
117
+ ident = {"id": "resolver", "client": hostname}
118
+
119
+ # Lazily initialize the resolver on first use to attach to the correct event loop.
120
+ if self.resolver is None:
121
+ loop = asyncio.get_running_loop()
122
+ self.resolver = aiodns.DNSResolver(loop=loop)
123
+
124
+ hostname_lower = hostname.lower()
125
+ # 1. Check hosts file cache
126
+ if ip := self.hosts_cache.get(hostname_lower):
127
+ logger.debug(
128
+ flm(
129
+ f"Resolved to {ip} from hosts file cache.",
130
+ ident,
131
+ self.verbose,
132
+ )
133
+ )
134
+ return [ip]
135
+
136
+ # 2. Query DNS using aiodns for IPv4 and IPv6 addresses concurrently
137
+ tasks = [
138
+ self.resolver.query(hostname, "A"),
139
+ self.resolver.query(hostname, "AAAA"),
140
+ ]
141
+ results = await asyncio.gather(*tasks, return_exceptions=True)
142
+
143
+ resolved_ips: set[str] = set()
144
+ for res in results:
145
+ if isinstance(res, list):
146
+ for record in res:
147
+ resolved_ips.add(record.host)
148
+ elif isinstance(res, aiodns.error.DNSError):
149
+ # Ignore common "not found" errors. Log other DNS errors for debugging.
150
+ if res.args[0] not in (
151
+ aiodns.error.ARES_ENODATA,
152
+ aiodns.error.ARES_ENOTFOUND,
153
+ ):
154
+ logger.debug(
155
+ flm(f"aiodns query failed: {res}", ident, self.verbose)
156
+ )
157
+ elif isinstance(res, Exception):
158
+ logger.warning(
159
+ flm(
160
+ f"Unexpected error during DNS resolution: {res}",
161
+ ident,
162
+ self.verbose,
163
+ )
164
+ )
165
+
166
+ if not resolved_ips:
167
+ raise OSError(f"Failed to resolve host: {hostname}")
168
+
169
+ return list(resolved_ips)
170
+
171
+
172
+ # Initialize the singleton instance so it's ready for other modules to import
173
+ resolver = Resolver.get_instance()
@@ -20,6 +20,7 @@ DEFAULT_ALLOWLIST: set[str] = {
20
20
  "data.jsdelivr.com", # jsDelivr API
21
21
  "esm.run", # jsDelivr for JavaScript modules
22
22
  "unpkg.com", # Unpkg CDN
23
+ "bing.com", # Bing for search functionality
23
24
  "twitter.com", # Twitter for social media integration
24
25
  "x.com", # X (formerly Twitter) for social media integration"
25
26
  }
@@ -79,12 +79,11 @@ async def handle_connection(
79
79
 
80
80
  # --- Authentication Check ---
81
81
  if auth_file_path:
82
- # Pass method and uri for Digest authentication calculation
82
+ # Pass method for Digest authentication calculation
83
83
  user_ident = await verify_credentials(
84
84
  client_reader,
85
85
  client_writer,
86
86
  method,
87
- uri,
88
87
  headers,
89
88
  auth_file_path,
90
89
  )
@@ -1,21 +1,26 @@
1
1
  import re
2
2
 
3
+ # Regex patterns for extracting host and port from a string
3
4
  REGEX_HOST = re.compile(r"(.+?):([0-9]{1,5})")
4
5
 
5
- REGEX_CONTENT_LENGTH = re.compile(
6
- r"\r\nContent-Length: ([0-9]+)\r\n", re.IGNORECASE
7
- )
8
-
9
6
 
10
7
  def get_host_and_port(
11
8
  hostname: str, default_port: str | None = None
12
9
  ) -> tuple[str, int]:
10
+ """Extracts the host and port from a hostname string."""
13
11
  if match := REGEX_HOST.search(hostname):
14
12
  return match.group(1), int(match.group(2))
15
13
  return hostname, int(default_port or "80")
16
14
 
17
15
 
16
+ # Regex pattern for extracting Content-Length from HTTP headers
17
+ REGEX_CONTENT_LENGTH = re.compile(
18
+ r"\r\nContent-Length: ([0-9]+)\r\n", re.IGNORECASE
19
+ )
20
+
21
+
18
22
  def get_content_length(header: str) -> int:
23
+ """Extracts the Content-Length from an HTTP header string."""
19
24
  if match := REGEX_CONTENT_LENGTH.search(header):
20
25
  return int(match.group(1))
21
26
  return 0
@@ -0,0 +1,3 @@
1
+ from importlib.metadata import version
2
+
3
+ VERSION: str = version("wormhole-proxy")
@@ -1,3 +0,0 @@
1
- from importlib.metadata import version
2
-
3
- VERSION = version("wormhole-proxy")
File without changes