REPOSITORY / ScuroNeko/mtg

Compare commits

DIFF REPOSITORY

Compare commits

..
130 Commits
Author SHA1 Message Date
9seconds 83a31e0458 Merge remote-tracking branch 'origin/stable' into v2 2026-04-07 18:10:41 +02:00
9seconds 95db895136 Merge remote-tracking branch 'origin/master' into stable 2026-04-07 18:10:26 +02:00
9seconds fb94d4a78b Update dependencies 2026-04-07 18:09:39 +02:00
Sergei ArkhipovandGitHub b5799500e4 Merge pull request #454 from 9seconds/tcp-notsent-lowat
Add TCP_NOTSENT_LOWAT setting
2026-04-07 15:45:53 +02:00
9seconds 437dacfaab Refactor socksopts per functionality, not per build flag 2026-04-07 15:42:22 +02:00
9seconds 0de8b28de8 Add TCP_NOTSENT_LOWAT setting 2026-04-07 15:39:26 +02:00
Sergei ArkhipovandGitHub 2f62e8055d Merge pull request #453 from 9seconds/tcp-user-timeout
Add TCP_USER_TIMEOUT support
2026-04-07 15:38:11 +02:00
9seconds b58ac669d5 Add TCP_USER_TIMEOUT support 2026-04-07 15:35:47 +02:00
Sergei ArkhipovandGitHub 60ab083def Merge pull request #452 from 9seconds/tcp-bbr
Use TCP BBR in a best-effort mode
2026-04-07 15:33:27 +02:00
9seconds 88f33debab Use TCP BBR in a best-effort mode
This small commits conditionally sets TCP BBR as a preferrable algorithm
for TCP congestion control
2026-04-07 14:54:56 +02:00
Sergei ArkhipovandGitHub ac40c94820 Merge pull request #451 from 9seconds/keepalive-config
Propagate keep alive settings from the config
2026-04-07 13:53:00 +02:00
9seconds 102f8a6cce Propagate keep alive settings from the config 2026-04-07 13:41:44 +02:00
Sergei ArkhipovandGitHub 83b43aecc1 Merge pull request #448 from 9seconds/no-default-tls-cipher
Do not use default TLS cipher
2026-04-07 12:13:21 +02:00
9seconds 5bf218f9ab Do not use default TLS cipher
As per RFC, if TLS server cannot pickup a suitable cipher from a client
list, it has to send handshake_failure alert. For us it means that we
have to route a request to a fronting domain, because we want to have it
exactly like a real webserver does. So, if it misbehaves, so do we.
2026-04-07 12:08:14 +02:00
Sergei ArkhipovandGitHub efc65f317e Merge pull request #447 from 9seconds/handshake-timeout
Add separate handshake timeout
2026-04-07 08:29:23 +02:00
9seconds 39ab5570a8 Small refactoring 2026-04-07 08:29:05 +02:00
9seconds eb564936c7 Add separate handshake timeout
This PR adds a new setting to the config: `network.timeout`. This setting
defines a time period during which all handshake procedures and
ceremonies must be completed. If not - connection is aborted. This
should help in situations when connection is established but client
cannot continue for some reason (for example, RST sent by some middle box).
2026-04-07 08:01:51 +02:00
Sergei ArkhipovandGitHub 74a81a986a Merge pull request #446 from runixer/fix/grease-cipher-suite
Fix DPI detection: replace GREASE cipher suite in FakeTLS ServerHello
2026-04-07 07:47:21 +02:00
Constantine bec321d190 Fix DPI detection: skip GREASE cipher suite in ClientHello parsing
Instead of echoing the first cipher suite from ClientHello (which is
often a GREASE value like 0x5a5a), iterate the list and pick the first
real cipher suite. This is what real TLS servers do per RFC 8701.

Production data shows two client profiles:
- 87% send GREASE first, then 0x1301 (TLS_AES_128_GCM_SHA256)
- 13% send 0xc02b first (TLS_ECDHE_ECDSA_WITH_AES_128_GCM_SHA256)

The fix correctly selects 0x1301 or 0xc02b respectively, matching
real server behavior. Fallback to 0x1301 if all suites are GREASE.

Add snapshot test with GREASE as first cipher suite.
2026-04-06 23:16:08 +03:00
Sergei ArkhipovandGitHub 45f958e527 Merge pull request #441 from appolimp/tcp-keepalive-idle-timeout
Improve TCP keepalive and idle timeout for mobile clients
2026-04-06 21:27:42 +02:00
appolimp 5f81ae3743 Improve TCP keepalive and idle timeout for mobile clients
TCP keepalive was configured (SetKeepAlivePeriod) but never actually
enabled (SO_KEEPALIVE) on accepted client connections. Go 1.26's
SetKeepAlivePeriod only sets TCP_KEEPIDLE — it does not call
setsockopt(SO_KEEPALIVE, 1). Without SO_KEEPALIVE the kernel never
sends probe packets, so dead connections from sleeping mobile clients
linger until the idle timeout fires.

Replace SetKeepAlive + SetKeepAlivePeriod with net.KeepAliveConfig
(available since Go 1.24) for explicit per-socket control:

  Idle:     30s   (time before first probe)
  Interval: 10s   (between probes)
  Count:    3     (failed probes to declare dead)

This detects dead connections in ~60s instead of relying on system
defaults (tcp_keepalive_intvl=75s, probes=9 → up to 11 minutes).

Increase the default idle timeout from 1 minute to 5 minutes.
MTProto clients send ping_delay_disconnect every ~60s, which resets
the idle timer. The previous 1-minute default created a race: if a
ping arrived even 1–2 seconds late the relay was killed. A 5-minute
window also survives typical mobile sleep periods (phone idle 2–5 min)
where the NAT mapping is still alive and the connection can resume
without reconnection.

Ref: #132
2026-04-04 12:01:33 +03:00
9seconds 1a81efcb6e Merge remote-tracking branch 'origin/stable' into v2 2026-04-01 17:17:22 +02:00
9seconds 2544c521ed Merge remote-tracking branch 'origin/master' into stable 2026-04-01 17:16:55 +02:00
9seconds 3a68ea5f2d Update goreleaser 2026-04-01 17:05:30 +02:00
Sergei ArkhipovandGitHub dbced77566 Merge pull request #433 from 9seconds/refactor-tls-fragmentation
Refactor TLS fragmenting
2026-04-01 14:30:21 +02:00
9seconds f4f969e702 Refactor TLS fragmenting 2026-04-01 14:01:24 +02:00
Sergei ArkhipovandGitHub e8368f7645 Merge pull request #431 from appolimp/tls-record-reassembly-pr
Support fragmented TLS handshake records
2026-04-01 09:33:13 +02:00
appolimp 38abee7d7f Support fragmented TLS handshake records
DPI bypass tools like ByeDPI fragment a single TLS record into multiple
records to evade censorship. This broke ReadClientHello because it
assumed the entire ClientHello arrives in one TLS record.

Add reassembleTLSHandshake that reads continuation records and
reconstructs a single TLS record before parsing and HMAC verification.
Per RFC 5246 Section 6.2.1, handshake messages may be fragmented
across multiple records — this is valid TLS behavior.
2026-04-01 09:05:24 +03:00
9seconds a3663fe8b5 Increase timeout for CI artifacts build 2026-03-31 22:13:42 +02:00
9seconds 2aa3321bd4 Add more forks 2026-03-31 19:03:48 +02:00
Sergei ArkhipovandGitHub 3793558c4c Merge pull request #430 from 9seconds/golang-idiomatic
More idiomatic Golang
2026-03-31 17:23:11 +02:00
9seconds b6427ee321 More idiomatic Golang 2026-03-31 15:07:01 +02:00
Sergei ArkhipovandGitHub 0c9fa5e710 Merge pull request #428 from 9seconds/auto-update-prio{
Change IP address set priority
2026-03-31 12:56:58 +02:00
9seconds 1fcec38aea Change IP address set priority
For a couple of releases we use collected IPs as a prioritized source
for connecting to Telegram. But apparently, they work way worse than it
should, and having connectivity to core ip ALWAYS gives better results.
Thus, this PR flips priorities, so users could have auto-update enabled
as a source of secondary addresses, not primary ones
2026-03-31 11:05:49 +02:00
Sergei ArkhipovandGitHub 89930631cf Merge pull request #426 from dolonet/fix/flaky-ci-race-and-bloom 2026-03-31 07:10:00 +02:00
dolonet eedee63143 Address review: use slices.Clone, simplify concurrent test
- Replace manual make+copy with slices.Clone in Snapshot()
- Remove redundant _ = len(data); Snapshot() call alone is
  sufficient to exercise the lock under -race
2026-03-30 16:17:51 +00:00
Alexey Dolotov 73c6a3aa37 fix: tighten ScoutConnCollected encapsulation and add concurrency test
- Move error check before Snapshot() to avoid unnecessary allocation
- Update existing tests to use Snapshot() instead of direct field access
- Add TestConcurrentAddSnapshot to explicitly exercise the mutex
2026-03-30 15:05:50 +03:00
Alexey Dolotov e54d9d60d3 fix: stabilize flaky CI tests
1. Add sync.Mutex to ScoutConnCollected to eliminate data race between
   Add()/MarkWrite() in readLoop and learn() iterating results.
   Introduce Snapshot() for safe read access.

2. Increase bloom filter test size from 500 to 100000 to prevent
   false negatives from random eviction in the stable bloom filter.

3. Use Require().NoError() in TestHTTPSRequest to prevent nil-pointer
   panic on resp.Body.Close() when the request fails.

Fixes #425
2026-03-30 14:50:32 +03:00
9seconds db2e6031a3 Merge remote-tracking branch 'origin/stable' into v2 2026-03-30 13:09:05 +02:00
9seconds 4b8da719ae Merge remote-tracking branch 'origin/master' into stable 2026-03-30 13:08:51 +02:00
9seconds a2de52f071 Update PGO 2026-03-30 13:08:32 +02:00
Sergei ArkhipovandGitHub b926a0590c Merge pull request #424 from dolonet/fix/relay-idle-timeout-shared-tracker
fix: use shared idle tracker for relay connections
2026-03-30 12:54:51 +02:00
Alexey Dolotov 58e8c8f982 ci: trigger tests 2026-03-30 13:51:17 +03:00
Alexey Dolotov 4642546b35 test: add idleTracker and connIdleTimeout tests
Cover shared idle tracker behavior:
- tracker lifecycle (new, idle after timeout, touch resets)
- read/write with data touches tracker
- read retries on timeout when tracker is not idle
- read closes on timeout when tracker is idle
- shared tracker prevents false timeout across directions
2026-03-30 13:21:51 +03:00
Alexey Dolotov 4627910238 fix: use shared idle tracker for relay connections
connIdleTimeout previously set per-direction deadlines independently.
During media downloads the client→telegram direction can be idle at the
application level while telegram→client is actively streaming data.
After IdleTimeout (default 1 min) the idle direction's ReadDeadline
fires, tearing down the entire relay and breaking media transfers.

Replace the per-direction timeout with a shared atomic timestamp that
both pump goroutines update on any successful Read or Write. When a
ReadDeadline fires on the idle direction, we check the shared tracker:
if the other direction was recently active, we retry instead of closing.
The connection is only torn down when both directions are idle for the
full timeout period.

This matches the documented IdleTimeout contract: "if we have any
message which will pass to either direction, a timer is reset."

Overhead: one atomic.Int64 (8 bytes) per connection pair, one
atomic.Store (~1 ns) per Read/Write with data, zero extra goroutines.

Fixes #423
2026-03-30 10:34:16 +03:00
9seconds 8b3c622ea6 Merge remote-tracking branch 'origin/stable' into v2 2026-03-29 23:33:24 +02:00
9seconds 9d43c2d759 Merge remote-tracking branch 'origin/master' into stable 2026-03-29 23:33:06 +02:00
9seconds 0840c7e3e5 Update PGO 2026-03-29 23:31:41 +02:00
9seconds de48e177b1 Update depndencies 2026-03-29 23:16:06 +02:00
Sergei ArkhipovandGitHub 3ed09146b9 Merge pull request #422 from 9seconds/release-ci
Build release artifacts in CI
2026-03-29 23:15:28 +02:00
9seconds 018bd2fdc1 Run release build 2026-03-29 22:54:08 +02:00
Sergei ArkhipovandGitHub 0ad3a06863 Merge pull request #421 from 9seconds/mips-save-mem
Decrease a relay buffer size for MIPS devices
2026-03-29 22:09:25 +02:00
9seconds d3a090d6b4 Decrease a relay buffer size for MIPS devices 2026-03-29 21:42:26 +02:00
Sergei ArkhipovandGitHub 1725a0d721 Merge pull request #420 from dolonet/fix/telegram-relay-idle-timeout
fix: apply idle timeout to Telegram relay
2026-03-29 17:14:58 +02:00
Alexey Dolotov 87988326ab retry CI 2026-03-29 17:02:53 +03:00
Alexey Dolotov ec271baab0 fix: apply idle timeout to Telegram relay
Wrap both sides of the Telegram relay in connIdleTimeout,
same as already done for domain fronting in #416.

Without this, if a client disappears (network drop, battery dies),
the TCP connection stays formally alive and the goroutine in the
worker pool blocks on io.CopyBuffer indefinitely. Under mass client
disconnects this accumulates zombie goroutines.

Fixes #417
2026-03-29 16:59:30 +03:00
Sergei ArkhipovandGitHub 139db15e83 Merge pull request #419 from 9seconds/timers
Remove clock goroutine
2026-03-29 15:46:55 +02:00
9seconds 0c030646f9 Remove clock goroutine
This is a followup for https://github.com/9seconds/mtg/issues/412 it
makes sense to manage timers inplace instead of creating for new
goroutines: saves memory
2026-03-29 15:39:33 +02:00
Sergei ArkhipovandGitHub 822560bede Merge pull request #418 from dolonet/public-ip-config
Add public-ipv4/public-ipv6 config options
2026-03-29 13:52:07 +02:00
Sergei ArkhipovandGitHub 9917f61bc8 Merge pull request #415 from dolonet/fix/event-stream-32bit-index-panic
fix: prevent index out of range panic on 32-bit platforms
2026-03-29 13:51:12 +02:00
Sergei ArkhipovandGitHub 735466b90d Merge pull request #416 from dolonet/fix/domain-fronting-idle-timeout
fix: apply idle timeout to domain fronting relay
2026-03-29 13:50:45 +02:00
Sergei ArkhipovandGitHub 6b51de8305 Merge pull request #414 from dolonet/optimize-per-connection-overhead
Reduce per-connection memory overhead
2026-03-29 13:49:08 +02:00
Alexey Dolotov 7b62e06e36 Retry CI: flaky antireplay bloom filter test 2026-03-29 00:53:20 +03:00
Alexey Dolotov 2b07c0037e Add public-ipv4/public-ipv6 config options for manual IP override
On some servers ifconfig.co is unreachable (e.g. Hetzner, AdGuard DNS
blocklists), causing 'mtg doctor' SNI-DNS check and 'mtg access' link
generation to fail. New config options allow specifying public IPs
manually, with automatic detection as fallback.

Fixes #405
2026-03-29 00:47:49 +03:00
Alexey Dolotov 450381ee16 ci: retrigger (flaky antireplay test) 2026-03-28 23:13:35 +03:00
Alexey Dolotov 46c33f1532 ci: retrigger (flaky antireplay test) 2026-03-28 23:04:33 +03:00
Alexey Dolotov f355512aa6 fix: address staticcheck lint issues
- avoid deprecated DefaultIdleTimeout, use time.Minute directly
- simplify embedded field selectors (QF1008)
2026-03-28 22:59:32 +03:00
Alexey Dolotov 289bb283b1 fix: close connection on worker pool overflow
When the worker pool rejected a connection (ErrPoolOverload), the
accepted net.Conn was never closed — leaking a file descriptor and
TCP socket per rejected connection. Under sustained traffic spikes this
compounds the problem: leaked descriptors reduce the capacity for new
dials (including to the fronting domain), accelerating the failure
cascade described in #378.
2026-03-28 22:52:39 +03:00
Alexey Dolotov 836090ebdf fix: apply idle timeout to domain fronting relay connections
Domain fronting relay (for non-Telegram traffic) had no idle timeout,
causing worker pool exhaustion under traffic spikes.

The ProxyOpts.IdleTimeout field existed but was never wired into the
proxy. Now domain fronting connections are wrapped with per-read/write
deadlines reset to the configured idle timeout (default 1m), so stale
or slowloris-style connections are reaped promptly.

Fixes #378
2026-03-28 22:47:39 +03:00
Alexey Dolotov 01402bdba2 fix: prevent index out of range panic on 32-bit platforms
On 32-bit architectures (e.g. ARM7), int is 32 bits wide.
Casting a uint32 hash value to int can overflow, producing a
negative number. Go's modulo operator preserves the sign, so
the channel index can become -1, causing a panic.

Perform the modulo in uint32 space before indexing to ensure
the result is always non-negative.

Fixes #413
2026-03-28 22:31:47 +03:00
Alexey Dolotov 026ec74dfd Reduce per-connection memory overhead
- Use sync.Pool for relay buffers instead of stack-allocated arrays.
  A [16379]byte on the goroutine stack forces Go to grow it to 32KB
  (next power of two). Pooled buffers keep goroutine stacks small.

- Same fix for doppelganger write buffer ([16384]byte in conn.start).

- Replace idle goroutines with context.AfterFunc in proxy.ServeConn
  and relay.Relay. These goroutines existed only to wait on ctx.Done()
  and close connections. AfterFunc achieves the same without allocating
  a goroutine until the context is actually cancelled.

Net effect: at 3000 concurrent connections on a 1-vCPU/961MB VPS,
the unmodified binary drops 246 connections and falls to 10 MB/s.
With these changes: zero failures, 63 MB/s, 31% lower RSS.

Closes #412
2026-03-28 13:24:39 +03:00
Sergei ArkhipovandGitHub cc4b6ce2f4 Merge pull request #409 from dolonet/cert-noise-calibration
Add dynamic cert noise calibration for FakeTLS handshake
2026-03-28 09:04:18 +01:00
Alexey Dolotov 9dfd992c1d Move cert noise calibration into doppelganger scout
Instead of a separate cert_probe.go that duplicates the scout's TLS
connection logic, measure the cert chain size directly from the same
HTTPS connections the scout already makes.

Changes:
- Extend ScoutConnResult with payloadLen field
- Add Write interception to ScoutConn for handshake boundary detection
- Scout.learn() now computes cert size (sum of ApplicationData between
  CCS and first client Write) alongside inter-record durations
- Ganger aggregates cert sizes across raids and exposes NoiseParams()
  via atomic pointer for lock-free reads from proxy goroutines
- Proxy reads NoiseParams from Ganger on each handshake instead of
  probing at startup
- Remove cert_probe.go, disk cache, and related config options
  (noise-cache-path, noise-cache-ttl, noise-probe-count)

Falls back to legacy 2500-4700 range until the first scout raid
completes (typically within 1-2 seconds of startup).
2026-03-27 16:34:42 +03:00
Alexey Dolotov 80213ad35d Add dynamic cert noise calibration for FakeTLS handshake
The hardcoded noise range (2500-4700 bytes) in the FakeTLS ServerHello
does not match the real certificate chain sizes of many popular fronting
domains (e.g., dl.google.com ≈ 6480 bytes, microsoft.com ≈ 13004 bytes).
This makes the proxy detectable by DPI systems that compare the
ApplicationData size with the real cert chain size for the SNI domain.

On startup, probe the fronting domain's actual TLS handshake size and
use the measured value ± jitter instead of the static range. Falls back
to the legacy 2500-4700 range if the probe fails.

Also adds optional caching of probe results between restarts
(noise-cache-path, noise-cache-ttl) and a configurable probe count
(noise-probe-count) under [defense.doppelganger].

Closes #408
2026-03-26 23:38:58 +03:00
Sergei ArkhipovandGitHub d32e8e8b97 Merge pull request #404 from 9seconds/codeql
Update stale codeql configuration
2026-03-25 10:19:15 +01:00
9seconds 60c57c2306 Update stale codeql configuration 2026-03-25 10:17:30 +01:00
Sergei ArkhipovandGitHub 0edd5e6f92 Merge pull request #402 from 9seconds/PGO
Update PGO
2026-03-24 21:07:10 +01:00
9seconds a8e4acb6f8 Update PGO 2026-03-24 21:04:30 +01:00
Sergei ArkhipovandGitHub 27d10e6820 Merge pull request #401 from 9seconds/fix-prof
Fix build with profiling
2026-03-24 20:50:41 +01:00
9seconds b47e13556e Fix build with profiling 2026-03-24 20:46:41 +01:00
9seconds de81ed565d Add mention of fork 2026-03-24 15:28:33 +01:00
9seconds 006fba1046 Merge remote-tracking branch 'origin/stable' into v2 2026-03-24 09:59:03 +01:00
9seconds 7b333ed833 Merge remote-tracking branch 'origin/master' into stable 2026-03-24 09:58:48 +01:00
9seconds 5adfee5dd4 Remove wrong binary 2026-03-24 09:58:13 +01:00
9seconds de89de2ad6 Merge remote-tracking branch 'origin/master' into stable 2026-03-24 09:57:43 +01:00
9seconds b0d37de0ec Update linter 2026-03-24 09:57:20 +01:00
9seconds 0cb25ba7ff Update go dependencies 2026-03-24 09:55:53 +01:00
9seconds 614acd7303 Mention doctor in README 2026-03-24 09:55:30 +01:00
Sergei ArkhipovandGitHub 4f5368aa2a Merge pull request #398 from 9seconds/docker-directory
Allow using directory bind mounts for a docker container
2026-03-24 09:00:59 +01:00
9seconds cfb5fe66be Allow using directory bind mounts for a docker container
This helps with a situation when some applications do not allow mounting
individual files, but whole directories. In that case users could mount
`/config` directory with a single file, `config.toml`: `-v
/path/to/dir:/config`. Also, there is a backward compatibility to using
a single `/config.toml`
2026-03-24 08:48:42 +01:00
Sergei ArkhipovandGitHub fb390d3417 Merge pull request #397 from 9seconds/doctor 2026-03-23 19:35:49 +01:00
9seconds f0ae4ce290 Validate domain fronting availability 2026-03-23 19:22:21 +01:00
9seconds b6b900e430 Refactoring 2026-03-23 19:12:14 +01:00
9seconds 8154f65e0e Add validation of telegram connectivity 2026-03-23 18:34:38 +01:00
9seconds a60523fed0 Add verification of time skewness 2026-03-23 15:28:48 +01:00
9seconds 63b147c287 Add doctor command for deprecated config values 2026-03-23 14:45:10 +01:00
Sergei ArkhipovandGitHub 21c0d18c7c Merge pull request #395 from roman901/master 2026-03-21 23:31:24 +01:00
Roman Shishkin 8f0bf47d56 Add Config.GetConcurrency with default fallback 2026-03-21 21:53:38 +03:00
9seconds 7fec30908a Merge remote-tracking branch 'origin/stable' into v2 2026-03-20 11:29:39 +01:00
9seconds 2eb0828f72 Merge remote-tracking branch 'origin/master' into stable 2026-03-20 11:29:24 +01:00
Sergei ArkhipovandGitHub d01e089f54 Merge pull request #386 from 9seconds/architectures
Add more architectures for mtg
2026-03-20 11:22:58 +01:00
Sergei ArkhipovandGitHub c736881792 Merge pull request #388 from 9seconds/doc-limits
Document a necessety of increasing limits for systemd unit
2026-03-20 11:16:06 +01:00
9seconds d5a118f125 Remove explicit pgo 2026-03-20 11:15:01 +01:00
9seconds d79a8f8406 Fix failed builds 2026-03-20 11:14:25 +01:00
9seconds 97932758d1 Add mips support 2026-03-20 11:14:25 +01:00
9seconds 1f7d1c0eea Add windows builds 2026-03-20 11:14:25 +01:00
9seconds 8c73dde928 Add build for AMD64v3 2026-03-20 11:14:25 +01:00
9seconds ded3fe26b9 Build for ARMv9 2026-03-20 11:14:25 +01:00
Sergei ArkhipovandGitHub 2f00adfe91 Merge pull request #385 from 9seconds/pgo
Add PGO
2026-03-20 11:14:00 +01:00
9seconds 049bee3d84 Document a necessety of increasing limits for systemd unit
It seems that default DynamicUser limits are very low. We have to
increase them anyway.
2026-03-20 11:13:03 +01:00
9seconds 4fbabfda2a Add PGO 2026-03-20 10:54:30 +01:00
9seconds fc72de9e39 Merge remote-tracking branch 'origin/stable' into v2 2026-03-19 18:52:35 +01:00
9seconds cb627f2a66 Merge remote-tracking branch 'origin/master' into stable 2026-03-19 18:52:11 +01:00
Sergei ArkhipovandGitHub 9ba6df0d1c Merge pull request #383 from 9seconds/avoid-double-buffering
Avoid double buffering in TLS hot path
2026-03-19 17:46:36 +01:00
9seconds 4a8d099aca Remove unused buffer 2026-03-19 17:39:57 +01:00
9seconds feb57004e1 Fix reslicing 2026-03-19 17:39:48 +01:00
9seconds cb436efd87 Avoid double buffering in TLS hot path 2026-03-19 17:37:51 +01:00
Sergei ArkhipovandGitHub 24148ea95c Merge pull request #382 from 9seconds/write-cond
Optimize waiting time for TLS chunker
2026-03-19 15:51:11 +01:00
9seconds 724904f50d Wait in doppel.Conn if there is anything to write 2026-03-19 15:42:00 +01:00
9seconds a23ae05f3b Remove SyncWrite 2026-03-19 13:47:08 +01:00
Sergei ArkhipovandGitHub b153a55149 Merge pull request #379 from 9seconds/fix-telegram-ips
Show ip of telegram endpoints in event stream
2026-03-18 22:46:47 +01:00
9seconds 913a38d13a Show real IP of the telegram endpoint in event stream 2026-03-18 22:05:34 +01:00
9seconds dc81f7981c Merge remote-tracking branch 'origin/stable' into v2 2026-03-16 23:56:10 +01:00
9seconds 9d5fd989e5 Merge remote-tracking branch 'origin/master' into stable 2026-03-16 23:55:56 +01:00
Sergei ArkhipovandGitHub 81703233b0 Merge pull request #368 from 9seconds/flake-tests
Fix flaky test
2026-03-16 23:55:01 +01:00
9seconds eb7720b11e Fix flaky test 2026-03-16 23:44:06 +01:00
Sergei ArkhipovandGitHub df7ddc3d6a Merge pull request #367 from saleacy/patch-1
fix: ensure network.Dial and MakeHTTPClient use socks5 proxy
2026-03-16 23:38:43 +01:00
saleacyandGitHub 3bc1e415f9 fix: ensure network.Dial and MakeHTTPClient use socks5 proxy
The package `network/v2/proxy_network.go` does not wrap `network.Dial`
and `network.MakeHTTPClient`, which causes them to bypass the SOCKS5
proxy and initiate TCP connections directly from the local machine.
2026-03-17 01:35:18 +08:00
Sergei ArkhipovandGitHub 306fa19ad6 Merge pull request #366 from Maks-2012/patch-1
Fix preferIPOnlyIPv6
2026-03-16 15:31:09 +01:00
Maks-2012andGitHub 079252d810 Fix preferIPOnlyIPv6 2026-03-16 16:10:38 +03:00
83 changed files with 2670 additions and 690 deletions
+3
View File
@@ -0,0 +1,3 @@
# git config merge.theirs.name "Always accept theirs"
# git config merge.theirs.driver "cp %B %A"
default.pgo binary merge=theirs
+33
View File
@@ -119,6 +119,39 @@ jobs:
- name: Run linter - name: Run linter
run: mise tasks run lint run: mise tasks run lint
artifacts:
name: Build release artifacts
runs-on: ubuntu-latest
timeout-minutes: 20
steps:
- name: Checkout
uses: actions/checkout@v6
with:
submodules: recursive
- uses: jdx/mise-action@v3
name: Install mise
- name: Cache Go modules
uses: actions/cache@v5
with:
path: ~/go/pkg/mod
key: ${{ runner.os }}-gomod-${{ hashFiles('go.sum') }}
restore-keys: |
${{ runner.os }}-gomod-
- name: Cache cross-compilation build
uses: actions/cache@v5
with:
path: ~/.cache/go-build
key: ${{ runner.os }}-goreleaser-${{ hashFiles('go.sum') }}-${{ hashFiles('**/*.go') }}
restore-keys: |
${{ runner.os }}-goreleaser-${{ hashFiles('go.sum') }}-
${{ runner.os }}-goreleaser-
- name: Run release
run: mise tasks run release
docker: docker:
name: Docker name: Docker
runs-on: ubuntu-latest runs-on: ubuntu-latest
+6 -4
View File
@@ -45,11 +45,13 @@ jobs:
steps: steps:
- name: Checkout repository - name: Checkout repository
uses: actions/checkout@v2 uses: actions/checkout@v6
with:
submodules: recursive
# Initializes the CodeQL tools for scanning. # Initializes the CodeQL tools for scanning.
- name: Initialize CodeQL - name: Initialize CodeQL
uses: github/codeql-action/init@v1 uses: github/codeql-action/init@v4
with: with:
languages: ${{ matrix.language }} languages: ${{ matrix.language }}
# If you wish to specify custom queries, you can do so here or in a config file. # If you wish to specify custom queries, you can do so here or in a config file.
@@ -60,7 +62,7 @@ jobs:
# Autobuild attempts to build any compiled languages (C/C++, C#, or Java). # Autobuild attempts to build any compiled languages (C/C++, C#, or Java).
# If this step fails, then you should remove it and run the build manually (see below) # If this step fails, then you should remove it and run the build manually (see below)
- name: Autobuild - name: Autobuild
uses: github/codeql-action/autobuild@v1 uses: github/codeql-action/autobuild@v4
# ️ Command-line programs to run using the OS shell. # ️ Command-line programs to run using the OS shell.
# 📚 https://git.io/JvXDl # 📚 https://git.io/JvXDl
@@ -74,4 +76,4 @@ jobs:
# make release # make release
- name: Perform CodeQL Analysis - name: Perform CodeQL Analysis
uses: github/codeql-action/analyze@v1 uses: github/codeql-action/analyze@v4
+81 -2
View File
@@ -10,13 +10,15 @@ before:
- go generate ./... - go generate ./...
builds: builds:
- binary: '{{ .ProjectName }}' - id: default
binary: '{{ .ProjectName }}'
goos: goos:
- darwin - darwin
- freebsd - freebsd
- linux - linux
- netbsd - netbsd
- openbsd - openbsd
- windows
goarch: goarch:
- 386 - 386
- amd64 - amd64
@@ -34,15 +36,92 @@ builds:
ignore: ignore:
- goos: darwin - goos: darwin
goarch: 386 goarch: 386
- goos: darwin
goarch: arm
- goos: freebsd - goos: freebsd
goarch: arm64 goarch: arm64
- goos: netbsd - goos: netbsd
goarch: arm64 goarch: arm64
- goos: openbsd - goos: openbsd
goarch: arm64 goarch: arm64
- goos: windows
goarch: 386
- goos: windows
goarch: arm
- id: mips
binary: '{{ .ProjectName }}'
goos:
- linux
goarch:
- mips
- mipsle
gomips:
- softfloat
env:
- CGO_ENABLED=0
flags:
- -trimpath
- -mod=readonly
ldflags: -s -w -X main.version={{ .Version }}
- id: arm64-v9
binary: '{{ .ProjectName }}'
goos:
- darwin
- linux
goarch:
- arm64
goarm64:
- v9.0
env:
- CGO_ENABLED=0
flags:
- -trimpath
- -mod=readonly
ldflags: -s -w -X main.version={{ .Version }}
- id: amd64-v3
binary: '{{ .ProjectName }}'
goos:
- darwin
- freebsd
- linux
- netbsd
- openbsd
- windows
goarch:
- amd64
goamd64:
- v3
env:
- CGO_ENABLED=0
flags:
- -trimpath
- -mod=readonly
ldflags: -s -w -X main.version={{ .Version }}
archives: archives:
- name_template: '{{ .ProjectName }}-{{ .Version }}-{{ .Os }}-{{ .Arch }}{{ if .Arm }}v{{ .Arm }}{{ end }}' - id: default
ids:
- default
- mips
name_template: '{{ .ProjectName }}-{{ .Version }}-{{ .Os }}-{{ .Arch }}{{ if .Arm }}v{{ .Arm }}{{ end }}'
formats:
- tar.gz
wrap_in_directory: true
format_overrides:
- goos: windows
formats:
- zip
files:
- LICENSE
- README.md
- SECURITY.md
- BEST_PRACTICES.md
- example.config.toml
- id: optimized
ids:
- arm64-v9
- amd64-v3
name_template: '{{ .ProjectName }}-{{ .Version }}-{{ .Os }}-{{ .Arch }}{{ if .Arm64 }}-{{ .Arm64 }}{{ end }}{{ if .Amd64 }}-{{ .Amd64 }}{{ end }}'
formats: formats:
- tar.gz - tar.gz
wrap_in_directory: true wrap_in_directory: true
+6
View File
@@ -16,6 +16,12 @@ sources = ["**/*.go", "go.mod", "go.sum"]
outputs = ["mtg"] outputs = ["mtg"]
run = "go build" run = "go build"
[tasks."build:prof"]
description = "Build binary with profiling enabled"
sources = ["**/*.go", "go.mod", "go.sum"]
outputs = ["mtg"]
run = "go build -tags prof"
[tasks.update] [tasks.update]
description = "Update dependencies" description = "Update dependencies"
run = [ run = [
+16 -2
View File
@@ -5,6 +5,19 @@ FROM golang:1.26-alpine AS build
ENV CGO_ENABLED=0 ENV CGO_ENABLED=0
# this is done for backward compatibility: before that we mounted a config
# into /config.toml. Some application allow mounting directories only,
# so it makes problems. So, instead we are going to do 2 steps:
# 1. Create /config/config.toml as a symlink to /config.toml
# 2. Force /mtg to use /config/config.toml
#
# it helps in both ways: users with directories could use /config directory
# and overlap a symlink by their bind mount. Old users could continue using
# /config.toml as a real config.
RUN set -x \
&& mkdir -p /config \
&& ln -sv /config.toml /config/config.toml
RUN --mount=type=cache,target=/var/cache/apk \ RUN --mount=type=cache,target=/var/cache/apk \
set -x \ set -x \
&& apk --update add \ && apk --update add \
@@ -20,7 +33,7 @@ RUN go mod download
COPY . /app COPY . /app
RUN set -x \ RUN set -x \
&& version="$(git describe --exact-match HEAD 2>/dev/null || git describe --tags --always)" \ && version="$(git describe --exact-match HEAD 2>/dev/null || git describe --tags --always 2>/dev/null || echo dev)" \
&& go build \ && go build \
-trimpath \ -trimpath \
-mod=readonly \ -mod=readonly \
@@ -35,8 +48,9 @@ RUN set -x \
FROM scratch FROM scratch
ENTRYPOINT ["/mtg"] ENTRYPOINT ["/mtg"]
CMD ["run", "/config.toml"] CMD ["run", "/config/config.toml"]
COPY --from=build /etc/ssl/certs/ca-certificates.crt /etc/ssl/certs/ca-certificates.crt COPY --from=build /etc/ssl/certs/ca-certificates.crt /etc/ssl/certs/ca-certificates.crt
COPY --from=build /app/mtg /mtg COPY --from=build /app/mtg /mtg
COPY --from=build /app/example.config.toml /config.toml COPY --from=build /app/example.config.toml /config.toml
COPY --from=build /config /config
+38
View File
@@ -29,6 +29,8 @@ are the most notable:
* [Official](https://github.com/TelegramMessenger/MTProxy) * [Official](https://github.com/TelegramMessenger/MTProxy)
* [Python](https://github.com/alexbers/mtprotoproxy) * [Python](https://github.com/alexbers/mtprotoproxy)
* [Erlang](https://github.com/seriyps/mtproto_proxy) * [Erlang](https://github.com/seriyps/mtproto_proxy)
* [Teleproxy (C)](https://github.com/teleproxy/teleproxy)
* [mtproto.zig (Zig)](https://github.com/sleep3r/mtproto.zig)
* [Telemt (Rust)](https://github.com/telemt/telemt) * [Telemt (Rust)](https://github.com/telemt/telemt)
You can use any of these. They work great and all implementations have You can use any of these. They work great and all implementations have
@@ -91,6 +93,9 @@ that probably matter.
software. I also believe that in the case of throwout proxies, this software. I also believe that in the case of throwout proxies, this
the feature is a useless luxury. the feature is a useless luxury.
This is very controversial topic. Please read [rationale (in russian)](https://github.com/9seconds/mtg/issues/376#issuecomment-4118726699)
and use [mtg-multi](https://github.com/dolonet/mtg-multi) fork if you are disagree with.
* **No adtag support** * **No adtag support**
Please read [Version 2](#version-2) chapter. Please read [Version 2](#version-2) chapter.
@@ -301,6 +306,38 @@ For example, you've bought a VPS from [Digital
Ocean](https://www.digitalocean.com/). Then it might be a good idea to Ocean](https://www.digitalocean.com/). Then it might be a good idea to
generate a secret for _digitalocean.com_ then. generate a secret for _digitalocean.com_ then.
### Check configuration
There is a special command for secret verification:
```
$ mtg doctor /path/to/my/config.toml
Deprecated options
✅ All good
Time skewness
✅ Time drift is -607.048µs, but tolerate-time-skewness is 5s
Validate native network connectivity
✅ DC 1
✅ DC 2
✅ DC 3
✅ DC 4
✅ DC 5
✅ DC 203
Validate network connectivity with proxy socks5://127.0.0.1:1080
✅ DC 1
✅ DC 2
✅ DC 3
✅ DC 4
✅ DC 5
✅ DC 203
Validate fronting domain connectivity
✅ xx.xx.xx.xx:yyy is reachable
Validate SNI-DNS match
✅ IP address xx.xx.xx.xx matches secret hostname <REDACTED>
```
It aims to find out possible inconsistencies and problems with your
configuration. It makes sense to run it before executing any relevant commands.
### Simple run mode ### Simple run mode
@@ -380,6 +417,7 @@ ExecStart=/usr/local/bin/mtg run /etc/mtg.toml
Restart=always Restart=always
RestartSec=3 RestartSec=3
DynamicUser=true DynamicUser=true
LimitNOFILE=65536
AmbientCapabilities=CAP_NET_BIND_SERVICE AmbientCapabilities=CAP_NET_BIND_SERVICE
[Install] [Install]
+1 -1
View File
@@ -12,7 +12,7 @@ type StableBloomFilterTestSuite struct {
} }
func (suite *StableBloomFilterTestSuite) TestOp() { func (suite *StableBloomFilterTestSuite) TestOp() {
filter := antireplay.NewStableBloomFilter(500, 0.001) filter := antireplay.NewStableBloomFilter(100000, 0.001)
suite.False(filter.SeenBefore([]byte{1, 2, 3})) suite.False(filter.SeenBefore([]byte{1, 2, 3}))
suite.False(filter.SeenBefore([]byte{4, 5, 6})) suite.False(filter.SeenBefore([]byte{4, 5, 6}))
BIN
View File
Binary file not shown.
+30
View File
@@ -0,0 +1,30 @@
package essentials
// TelegramCoreAddresses are publicly known addresses of Telegram core network.
var TelegramCoreAddresses = map[int][]string{
1: {
"149.154.175.50:443",
"[2001:b28:f23d:f001::a]:443",
},
2: {
"149.154.167.51:443",
"95.161.76.100:443",
"[2001:67c:04e8:f002::a]:443",
},
3: {
"149.154.175.100:443",
"[2001:b28:f23d:f003::a]:443",
},
4: {
"149.154.167.91:443",
"[2001:67c:04e8:f004::a]:443",
},
5: {
"149.154.171.5:443",
"[2001:b28:f23f:f005::a]:443",
},
203: {
"91.105.192.100:443",
"[2a0a:f280:0203:000a:5000:0000:0000:0100]:443",
},
}
+1 -1
View File
@@ -38,7 +38,7 @@ func (e EventStream) Send(ctx context.Context, evt mtglib.Event) {
select { select {
case <-ctx.Done(): case <-ctx.Done():
case <-e.ctx.Done(): case <-e.ctx.Done():
case e.chans[int(chanNo)%len(e.chans)] <- evt: case e.chans[chanNo%uint32(len(e.chans))] <- evt:
} }
} }
+20 -5
View File
@@ -48,6 +48,13 @@ concurrency = 8192
# Only ipv4 connectivity is used # Only ipv4 connectivity is used
prefer-ip = "prefer-ipv6" prefer-ip = "prefer-ipv6"
# Public IP addresses of this server. Used by 'mtg access' to generate
# proxy links and by 'mtg doctor' to validate SNI-DNS match.
# If not set, mtg tries to detect them automatically via ifconfig.co.
# Set these if ifconfig.co is unreachable from your server.
# public-ipv4 = "1.2.3.4"
# public-ipv6 = "2001:db8::1"
# If this setting is set, then mtg will try to get proxy updates from Telegram # If this setting is set, then mtg will try to get proxy updates from Telegram
# Usually this is completely fine to have it disabled, because mtg has a list # Usually this is completely fine to have it disabled, because mtg has a list
# of some core proxies hardcoded. # of some core proxies hardcoded.
@@ -197,14 +204,22 @@ proxies = [
# define a global timeout on establishing of network connections. idle # define a global timeout on establishing of network connections. idle
# means a timeout on pumping data between sockset when nothing is # means a timeout on pumping data between sockset when nothing is
# happening. # happening.
#
# please be noticed that handshakes have no timeouts intentionally. You can
# find a reasoning here:
# https://www.ndss-symposium.org/wp-content/uploads/2020/02/23087-paper.pdf
[network.timeout] [network.timeout]
tcp = "5s" tcp = "5s"
http = "10s" http = "10s"
idle = "1m" idle = "5m"
handshake = "10s"
# this defines a configuration for TCP keep alives. Default values are taken
# from Golang default behavior.
[network.keep-alive]
disabled = false
# idle means a time period after which we start sending TCP Keep Alive probes
idle = "15s"
# interval is a period between 2 consecutive probes
interval = "15s"
# if we miss that many probes, a connection will be considered as a dead one.
count = 9
# mtg has to mimic real websites. It does not mean domain fronting, it also # mtg has to mimic real websites. It does not mean domain fronting, it also
# means that traffic characteristics should be similar to real world traffic. # means that traffic characteristics should be similar to real world traffic.
+6 -5
View File
@@ -4,18 +4,18 @@ go 1.26
require ( require (
github.com/OneOfOne/xxhash v1.2.8 github.com/OneOfOne/xxhash v1.2.8
github.com/alecthomas/kong v1.14.0 github.com/alecthomas/kong v1.15.0
github.com/alecthomas/units v0.0.0-20240927000941-0f3dac36c52b github.com/alecthomas/units v0.0.0-20240927000941-0f3dac36c52b
github.com/armon/go-socks5 v0.0.0-20160902184237-e75332964ef5 github.com/armon/go-socks5 v0.0.0-20160902184237-e75332964ef5
github.com/babolivier/go-doh-client v0.0.0-20201028162107-a76cff4cb8b6 github.com/babolivier/go-doh-client v0.0.0-20201028162107-a76cff4cb8b6
github.com/d4l3k/messagediff v1.2.1 // indirect github.com/d4l3k/messagediff v1.2.1 // indirect
github.com/jarcoal/httpmock v1.0.8 github.com/jarcoal/httpmock v1.0.8
github.com/mccutchen/go-httpbin v1.1.1 github.com/mccutchen/go-httpbin v1.1.1
github.com/panjf2000/ants/v2 v2.11.6 github.com/panjf2000/ants/v2 v2.12.0
github.com/prometheus/client_golang v1.23.2 github.com/prometheus/client_golang v1.23.2
github.com/prometheus/common v0.67.5 // indirect github.com/prometheus/common v0.67.5 // indirect
github.com/prometheus/procfs v0.20.1 // indirect github.com/prometheus/procfs v0.20.1 // indirect
github.com/rs/zerolog v1.34.0 github.com/rs/zerolog v1.35.0
github.com/smira/go-statsd v1.3.4 github.com/smira/go-statsd v1.3.4
github.com/stretchr/objx v0.5.2 // indirect github.com/stretchr/objx v0.5.2 // indirect
github.com/stretchr/testify v1.11.1 github.com/stretchr/testify v1.11.1
@@ -27,8 +27,9 @@ require (
) )
require ( require (
github.com/ncruces/go-dns v1.3.2 github.com/beevik/ntp v1.5.0
github.com/pelletier/go-toml/v2 v2.2.4 github.com/ncruces/go-dns v1.3.3
github.com/pelletier/go-toml/v2 v2.3.0
github.com/pires/go-proxyproto v0.11.0 github.com/pires/go-proxyproto v0.11.0
github.com/things-go/go-socks5 v0.1.0 github.com/things-go/go-socks5 v0.1.0
github.com/txthinking/socks5 v0.0.0-20251011041537-5c31f201a10e github.com/txthinking/socks5 v0.0.0-20251011041537-5c31f201a10e
+12 -19
View File
@@ -2,8 +2,8 @@ github.com/OneOfOne/xxhash v1.2.8 h1:31czK/TI9sNkxIKfaUfGlU47BAxQ0ztGgd9vPyqimf8
github.com/OneOfOne/xxhash v1.2.8/go.mod h1:eZbhyaAYD41SGSSsnmcpxVoRiQ/MPUTjUdIIOT9Um7Q= github.com/OneOfOne/xxhash v1.2.8/go.mod h1:eZbhyaAYD41SGSSsnmcpxVoRiQ/MPUTjUdIIOT9Um7Q=
github.com/alecthomas/assert/v2 v2.11.0 h1:2Q9r3ki8+JYXvGsDyBXwH3LcJ+WK5D0gc5E8vS6K3D0= github.com/alecthomas/assert/v2 v2.11.0 h1:2Q9r3ki8+JYXvGsDyBXwH3LcJ+WK5D0gc5E8vS6K3D0=
github.com/alecthomas/assert/v2 v2.11.0/go.mod h1:Bze95FyfUr7x34QZrjL+XP+0qgp/zg8yS+TtBj1WA3k= github.com/alecthomas/assert/v2 v2.11.0/go.mod h1:Bze95FyfUr7x34QZrjL+XP+0qgp/zg8yS+TtBj1WA3k=
github.com/alecthomas/kong v1.14.0 h1:gFgEUZWu2ZmZ+UhyZ1bDhuutbKN1nTtJTwh19Wsn21s= github.com/alecthomas/kong v1.15.0 h1:BVJstKbpO73zKpmIu+m/aLRrNmWwxXPIGTNin9VmLVI=
github.com/alecthomas/kong v1.14.0/go.mod h1:wrlbXem1CWqUV5Vbmss5ISYhsVPkBb1Yo7YKJghju2I= github.com/alecthomas/kong v1.15.0/go.mod h1:wrlbXem1CWqUV5Vbmss5ISYhsVPkBb1Yo7YKJghju2I=
github.com/alecthomas/repr v0.5.2 h1:SU73FTI9D1P5UNtvseffFSGmdNci/O6RsqzeXJtP0Qs= github.com/alecthomas/repr v0.5.2 h1:SU73FTI9D1P5UNtvseffFSGmdNci/O6RsqzeXJtP0Qs=
github.com/alecthomas/repr v0.5.2/go.mod h1:Fr0507jx4eOXV7AlPV6AVZLYrLIuIeSOWtW57eE/O/4= github.com/alecthomas/repr v0.5.2/go.mod h1:Fr0507jx4eOXV7AlPV6AVZLYrLIuIeSOWtW57eE/O/4=
github.com/alecthomas/units v0.0.0-20240927000941-0f3dac36c52b h1:mimo19zliBX/vSQ6PWWSL9lK8qwHozUj03+zLoEB8O0= github.com/alecthomas/units v0.0.0-20240927000941-0f3dac36c52b h1:mimo19zliBX/vSQ6PWWSL9lK8qwHozUj03+zLoEB8O0=
@@ -12,18 +12,18 @@ github.com/armon/go-socks5 v0.0.0-20160902184237-e75332964ef5 h1:0CwZNZbxp69SHPd
github.com/armon/go-socks5 v0.0.0-20160902184237-e75332964ef5/go.mod h1:wHh0iHkYZB8zMSxRWpUBQtwG5a7fFgvEO+odwuTv2gs= github.com/armon/go-socks5 v0.0.0-20160902184237-e75332964ef5/go.mod h1:wHh0iHkYZB8zMSxRWpUBQtwG5a7fFgvEO+odwuTv2gs=
github.com/babolivier/go-doh-client v0.0.0-20201028162107-a76cff4cb8b6 h1:4NNbNM2Iq/k57qEu7WfL67UrbPq1uFWxW4qODCohi+0= github.com/babolivier/go-doh-client v0.0.0-20201028162107-a76cff4cb8b6 h1:4NNbNM2Iq/k57qEu7WfL67UrbPq1uFWxW4qODCohi+0=
github.com/babolivier/go-doh-client v0.0.0-20201028162107-a76cff4cb8b6/go.mod h1:J29hk+f9lJrblVIfiJOtTFk+OblBawmib4uz/VdKzlg= github.com/babolivier/go-doh-client v0.0.0-20201028162107-a76cff4cb8b6/go.mod h1:J29hk+f9lJrblVIfiJOtTFk+OblBawmib4uz/VdKzlg=
github.com/beevik/ntp v1.5.0 h1:y+uj/JjNwlY2JahivxYvtmv4ehfi3h74fAuABB9ZSM4=
github.com/beevik/ntp v1.5.0/go.mod h1:mJEhBrwT76w9D+IfOEGvuzyuudiW9E52U2BaTrMOYow=
github.com/beorn7/perks v1.0.1 h1:VlbKKnNfV8bJzeqoa4cOKqO6bYr3WgKZxO8Z16+hsOM= github.com/beorn7/perks v1.0.1 h1:VlbKKnNfV8bJzeqoa4cOKqO6bYr3WgKZxO8Z16+hsOM=
github.com/beorn7/perks v1.0.1/go.mod h1:G2ZrVWU2WbWT9wwq4/hrbKbnv/1ERSJQ0ibhJ6rlkpw= github.com/beorn7/perks v1.0.1/go.mod h1:G2ZrVWU2WbWT9wwq4/hrbKbnv/1ERSJQ0ibhJ6rlkpw=
github.com/cespare/xxhash/v2 v2.3.0 h1:UL815xU9SqsFlibzuggzjXhog7bL6oX9BbNZnL2UFvs= github.com/cespare/xxhash/v2 v2.3.0 h1:UL815xU9SqsFlibzuggzjXhog7bL6oX9BbNZnL2UFvs=
github.com/cespare/xxhash/v2 v2.3.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs= github.com/cespare/xxhash/v2 v2.3.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs=
github.com/coreos/go-systemd/v22 v22.5.0/go.mod h1:Y58oyj3AT4RCenI/lSvhwexgC+NSVTIJ3seZv2GcEnc=
github.com/creack/pty v1.1.9/go.mod h1:oKZEueFk5CKHvIhNR5MUki03XCEU+Q6VDXinZuGJ33E= github.com/creack/pty v1.1.9/go.mod h1:oKZEueFk5CKHvIhNR5MUki03XCEU+Q6VDXinZuGJ33E=
github.com/d4l3k/messagediff v1.2.1 h1:ZcAIMYsUg0EAp9X+tt8/enBE/Q8Yd5kzPynLyKptt9U= github.com/d4l3k/messagediff v1.2.1 h1:ZcAIMYsUg0EAp9X+tt8/enBE/Q8Yd5kzPynLyKptt9U=
github.com/d4l3k/messagediff v1.2.1/go.mod h1:Oozbb1TVXFac9FtSIxHBMnBCq2qeH/2KkEQxENCrlLo= github.com/d4l3k/messagediff v1.2.1/go.mod h1:Oozbb1TVXFac9FtSIxHBMnBCq2qeH/2KkEQxENCrlLo=
github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c= github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c=
github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
github.com/godbus/dbus/v5 v5.0.4/go.mod h1:xhWf0FNVPg57R7Z0UbKHbJfkEywrmjJnf7w5xrFpKfA=
github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8= github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8=
github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU= github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU=
github.com/hexops/gotextdiff v1.0.3 h1:gitA9+qJrrTCsiCl7+kh75nPqQt1cx4ZkudSTLoUqJM= github.com/hexops/gotextdiff v1.0.3 h1:gitA9+qJrrTCsiCl7+kh75nPqQt1cx4ZkudSTLoUqJM=
@@ -38,11 +38,8 @@ github.com/kr/text v0.2.0 h1:5Nx0Ya0ZqY2ygV366QzturHI13Jq95ApcVaJBhpS+AY=
github.com/kr/text v0.2.0/go.mod h1:eLer722TekiGuMkidMxC/pM04lWEeraHUUmBw8l2grE= github.com/kr/text v0.2.0/go.mod h1:eLer722TekiGuMkidMxC/pM04lWEeraHUUmBw8l2grE=
github.com/kylelemons/godebug v1.1.0 h1:RPNrshWIDI6G2gRW9EHilWtl7Z6Sb1BR0xunSBf0SNc= github.com/kylelemons/godebug v1.1.0 h1:RPNrshWIDI6G2gRW9EHilWtl7Z6Sb1BR0xunSBf0SNc=
github.com/kylelemons/godebug v1.1.0/go.mod h1:9/0rRGxNHcop5bhtWyNeEfOS8JIWk580+fNqagV/RAw= github.com/kylelemons/godebug v1.1.0/go.mod h1:9/0rRGxNHcop5bhtWyNeEfOS8JIWk580+fNqagV/RAw=
github.com/mattn/go-colorable v0.1.13/go.mod h1:7S9/ev0klgBDR4GtXTXX8a3vIGJpMovkB8vQcUbaXHg=
github.com/mattn/go-colorable v0.1.14 h1:9A9LHSqF/7dyVVX6g0U9cwm9pG3kP9gSzcuIPHPsaIE= github.com/mattn/go-colorable v0.1.14 h1:9A9LHSqF/7dyVVX6g0U9cwm9pG3kP9gSzcuIPHPsaIE=
github.com/mattn/go-colorable v0.1.14/go.mod h1:6LmQG8QLFO4G5z1gPvYEzlUgJ2wF+stgPZH1UqBm1s8= github.com/mattn/go-colorable v0.1.14/go.mod h1:6LmQG8QLFO4G5z1gPvYEzlUgJ2wF+stgPZH1UqBm1s8=
github.com/mattn/go-isatty v0.0.16/go.mod h1:kYGgaQfpe5nmfYZH+SKPsOc2e4SrIfOl2e/yFXSvRLM=
github.com/mattn/go-isatty v0.0.19/go.mod h1:W+V8PltTTMOvKvAeJH7IuucS94S2C6jfK/D7dTCTo3Y=
github.com/mattn/go-isatty v0.0.20 h1:xfD0iDuEKnDkl03q4limB+vH+GxLEtL/jb4xVJSWWEY= github.com/mattn/go-isatty v0.0.20 h1:xfD0iDuEKnDkl03q4limB+vH+GxLEtL/jb4xVJSWWEY=
github.com/mattn/go-isatty v0.0.20/go.mod h1:W+V8PltTTMOvKvAeJH7IuucS94S2C6jfK/D7dTCTo3Y= github.com/mattn/go-isatty v0.0.20/go.mod h1:W+V8PltTTMOvKvAeJH7IuucS94S2C6jfK/D7dTCTo3Y=
github.com/mccutchen/go-httpbin v1.1.1 h1:aEws49HEJEyXHLDnshQVswfUlCVoS8g6h9YaDyaW7RE= github.com/mccutchen/go-httpbin v1.1.1 h1:aEws49HEJEyXHLDnshQVswfUlCVoS8g6h9YaDyaW7RE=
@@ -51,17 +48,16 @@ github.com/miekg/dns v1.1.51 h1:0+Xg7vObnhrz/4ZCZcZh7zPXlmU0aveS2HDBd0m0qSo=
github.com/miekg/dns v1.1.51/go.mod h1:2Z9d3CP1LQWihRZUf29mQ19yDThaI4DAYzte2CaQW5c= github.com/miekg/dns v1.1.51/go.mod h1:2Z9d3CP1LQWihRZUf29mQ19yDThaI4DAYzte2CaQW5c=
github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822 h1:C3w9PqII01/Oq1c1nUAm88MOHcQC9l5mIlSMApZMrHA= github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822 h1:C3w9PqII01/Oq1c1nUAm88MOHcQC9l5mIlSMApZMrHA=
github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822/go.mod h1:+n7T8mK8HuQTcFwEeznm/DIxMOiR9yIdICNftLE1DvQ= github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822/go.mod h1:+n7T8mK8HuQTcFwEeznm/DIxMOiR9yIdICNftLE1DvQ=
github.com/ncruces/go-dns v1.3.2 h1:kBLuUZBgkQ4qF4WDXZRQ4rG0Gk6sLVJQ5tESkWrxUa0= github.com/ncruces/go-dns v1.3.3 h1:59OV7XoJrTCoUMZjWRVs4GOjtntMTZqiQ5Mn+BT13hk=
github.com/ncruces/go-dns v1.3.2/go.mod h1:tuzixNY8PY/M7yUzcvRbUaeLs3ifIdydpi5H2bfRU+s= github.com/ncruces/go-dns v1.3.3/go.mod h1:tuzixNY8PY/M7yUzcvRbUaeLs3ifIdydpi5H2bfRU+s=
github.com/panjf2000/ants/v2 v2.11.6 h1:JKsoIUukIoCO0sP0gcOqdyoXmpyKXuU6fC57rODtpug= github.com/panjf2000/ants/v2 v2.12.0 h1:u9JhESo83i/GkZnhfTNuFMMWcNt7mnV1bGJ6FT4wXH8=
github.com/panjf2000/ants/v2 v2.11.6/go.mod h1:8u92CYMUc6gyvTIw8Ru7Mt7+/ESnJahz5EVtqfrilek= github.com/panjf2000/ants/v2 v2.12.0/go.mod h1:tSQuaNQ6r6NRhPt+IZVUevvDyFMTs+eS4ztZc52uJTY=
github.com/patrickmn/go-cache v2.1.0+incompatible h1:HRMgzkcYKYpi3C8ajMPV8OFXaaRUnok+kx1WdO15EQc= github.com/patrickmn/go-cache v2.1.0+incompatible h1:HRMgzkcYKYpi3C8ajMPV8OFXaaRUnok+kx1WdO15EQc=
github.com/patrickmn/go-cache v2.1.0+incompatible/go.mod h1:3Qf8kWWT7OJRJbdiICTKqZju1ZixQ/KpMGzzAfe6+WQ= github.com/patrickmn/go-cache v2.1.0+incompatible/go.mod h1:3Qf8kWWT7OJRJbdiICTKqZju1ZixQ/KpMGzzAfe6+WQ=
github.com/pelletier/go-toml/v2 v2.2.4 h1:mye9XuhQ6gvn5h28+VilKrrPoQVanw5PMw/TB0t5Ec4= github.com/pelletier/go-toml/v2 v2.3.0 h1:k59bC/lIZREW0/iVaQR8nDHxVq8OVlIzYCOJf421CaM=
github.com/pelletier/go-toml/v2 v2.2.4/go.mod h1:2gIqNv+qfxSVS7cM2xJQKtLSTLUE9V8t9Stt+h56mCY= github.com/pelletier/go-toml/v2 v2.3.0/go.mod h1:2gIqNv+qfxSVS7cM2xJQKtLSTLUE9V8t9Stt+h56mCY=
github.com/pires/go-proxyproto v0.11.0 h1:gUQpS85X/VJMdUsYyEgyn59uLJvGqPhJV5YvG68wXH4= github.com/pires/go-proxyproto v0.11.0 h1:gUQpS85X/VJMdUsYyEgyn59uLJvGqPhJV5YvG68wXH4=
github.com/pires/go-proxyproto v0.11.0/go.mod h1:ZKAAyp3cgy5Y5Mo4n9AlScrkCZwUy0g3Jf+slqQVcuU= github.com/pires/go-proxyproto v0.11.0/go.mod h1:ZKAAyp3cgy5Y5Mo4n9AlScrkCZwUy0g3Jf+slqQVcuU=
github.com/pkg/errors v0.9.1/go.mod h1:bwawxfHBFNV+L2hUp1rHADufV3IMtnDRdf1r5NINEl0=
github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM= github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM=
github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4=
github.com/prometheus/client_golang v1.23.2 h1:Je96obch5RDVy3FDMndoUsjAhG5Edi49h0RJWRi/o0o= github.com/prometheus/client_golang v1.23.2 h1:Je96obch5RDVy3FDMndoUsjAhG5Edi49h0RJWRi/o0o=
@@ -74,9 +70,8 @@ github.com/prometheus/procfs v0.20.1 h1:XwbrGOIplXW/AU3YhIhLODXMJYyC1isLFfYCsTEy
github.com/prometheus/procfs v0.20.1/go.mod h1:o9EMBZGRyvDrSPH1RqdxhojkuXstoe4UlK79eF5TGGo= github.com/prometheus/procfs v0.20.1/go.mod h1:o9EMBZGRyvDrSPH1RqdxhojkuXstoe4UlK79eF5TGGo=
github.com/rogpeppe/go-internal v1.14.1 h1:UQB4HGPB6osV0SQTLymcB4TgvyWu6ZyliaW0tI/otEQ= github.com/rogpeppe/go-internal v1.14.1 h1:UQB4HGPB6osV0SQTLymcB4TgvyWu6ZyliaW0tI/otEQ=
github.com/rogpeppe/go-internal v1.14.1/go.mod h1:MaRKkUm5W0goXpeCfT7UZI6fk/L7L7so1lCWt35ZSgc= github.com/rogpeppe/go-internal v1.14.1/go.mod h1:MaRKkUm5W0goXpeCfT7UZI6fk/L7L7so1lCWt35ZSgc=
github.com/rs/xid v1.6.0/go.mod h1:7XoLgs4eV+QndskICGsho+ADou8ySMSjJKDIan90Nz0= github.com/rs/zerolog v1.35.0 h1:VD0ykx7HMiMJytqINBsKcbLS+BJ4WYjz+05us+LRTdI=
github.com/rs/zerolog v1.34.0 h1:k43nTLIwcTVQAncfCw4KZ2VY6ukYoZaBPNOE8txlOeY= github.com/rs/zerolog v1.35.0/go.mod h1:EjML9kdfa/RMA7h/6z6pYmq1ykOuA8/mjWaEvGI+jcw=
github.com/rs/zerolog v1.34.0/go.mod h1:bJsvje4Z08ROH4Nhs5iH600c3IkWhwp44iRc54W6wYQ=
github.com/smira/go-statsd v1.3.4 h1:kBYWcLSGT+qC6JVbvfz48kX7mQys32fjDOPrfmsSx2c= github.com/smira/go-statsd v1.3.4 h1:kBYWcLSGT+qC6JVbvfz48kX7mQys32fjDOPrfmsSx2c=
github.com/smira/go-statsd v1.3.4/go.mod h1:RjdsESPgDODtg1VpVVf9MJrEW2Hw0wtRNbmB1CAhu6A= github.com/smira/go-statsd v1.3.4/go.mod h1:RjdsESPgDODtg1VpVVf9MJrEW2Hw0wtRNbmB1CAhu6A=
github.com/stretchr/objx v0.1.0/go.mod h1:HFkY916IF+rwdDfMAkV7OtwuqBVzrE8GR6GFx+wExME= github.com/stretchr/objx v0.1.0/go.mod h1:HFkY916IF+rwdDfMAkV7OtwuqBVzrE8GR6GFx+wExME=
@@ -131,10 +126,8 @@ golang.org/x/sys v0.0.0-20201119102817-f84b799fce68/go.mod h1:h1NjWce9XRLGQEsW7w
golang.org/x/sys v0.0.0-20210615035016-665e8c7367d1/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.0.0-20210615035016-665e8c7367d1/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.0.0-20220520151302-bc2c85ada10a/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.0.0-20220520151302-bc2c85ada10a/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.0.0-20220722155257-8c9f86f7a55f/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.0.0-20220722155257-8c9f86f7a55f/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.0.0-20220811171246-fbc7d0a398ab/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.2.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.2.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.6.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.6.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.12.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.42.0 h1:omrd2nAlyT5ESRdCLYdm3+fMfNFE/+Rf4bDIQImRJeo= golang.org/x/sys v0.42.0 h1:omrd2nAlyT5ESRdCLYdm3+fMfNFE/+Rf4bDIQImRJeo=
golang.org/x/sys v0.42.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw= golang.org/x/sys v0.42.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
golang.org/x/term v0.0.0-20201126162022-7de9c90e9dd1/go.mod h1:bj7SfCRtBDWHUb9snDiAeCFNEtKQo2Wmx5Cou7ajbmo= golang.org/x/term v0.0.0-20201126162022-7de9c90e9dd1/go.mod h1:bj7SfCRtBDWHUb9snDiAeCFNEtKQo2Wmx5Cou7ajbmo=
+8 -47
View File
@@ -1,22 +1,16 @@
package cli package cli
import ( import (
"context"
"encoding/json" "encoding/json"
"fmt" "fmt"
"io"
"net" "net"
"net/http"
"net/url" "net/url"
"os" "os"
"strconv" "strconv"
"strings"
"sync" "sync"
"github.com/9seconds/mtg/v2/essentials"
"github.com/9seconds/mtg/v2/internal/config" "github.com/9seconds/mtg/v2/internal/config"
"github.com/9seconds/mtg/v2/internal/utils" "github.com/9seconds/mtg/v2/internal/utils"
"github.com/9seconds/mtg/v2/mtglib"
) )
type accessResponse struct { type accessResponse struct {
@@ -65,7 +59,10 @@ func (a *Access) Run(cli *CLI, version string) error {
wg.Go(func() { wg.Go(func() {
ip := a.PublicIPv4 ip := a.PublicIPv4
if ip == nil { if ip == nil {
ip = a.getIP(ntw, "tcp4") ip = conf.PublicIPv4.Get(nil)
}
if ip == nil {
ip = getIP(ntw, "tcp4")
} }
if ip != nil { if ip != nil {
@@ -77,7 +74,10 @@ func (a *Access) Run(cli *CLI, version string) error {
wg.Go(func() { wg.Go(func() {
ip := a.PublicIPv6 ip := a.PublicIPv6
if ip == nil { if ip == nil {
ip = a.getIP(ntw, "tcp6") ip = conf.PublicIPv6.Get(nil)
}
if ip == nil {
ip = getIP(ntw, "tcp6")
} }
if ip != nil { if ip != nil {
@@ -100,45 +100,6 @@ func (a *Access) Run(cli *CLI, version string) error {
return nil return nil
} }
func (a *Access) getIP(ntw mtglib.Network, protocol string) net.IP {
dialer := ntw.NativeDialer()
client := ntw.MakeHTTPClient(func(ctx context.Context, network, address string) (essentials.Conn, error) {
conn, err := dialer.DialContext(ctx, protocol, address)
if err != nil {
return nil, err
}
return essentials.WrapNetConn(conn), err
})
req, err := http.NewRequest(http.MethodGet, "https://ifconfig.co", nil) //nolint: noctx
if err != nil {
panic(err)
}
req.Header.Add("Accept", "text/plain")
resp, err := client.Do(req)
if err != nil {
return nil
}
if resp.StatusCode != http.StatusOK {
return nil
}
defer func() {
io.Copy(io.Discard, resp.Body) //nolint: errcheck
resp.Body.Close() //nolint: errcheck
}()
data, err := io.ReadAll(resp.Body)
if err != nil {
return nil
}
return net.ParseIP(strings.TrimSpace(string(data)))
}
func (a *Access) makeURLs(conf *config.Config, ip net.IP) *accessResponseURLs { func (a *Access) makeURLs(conf *config.Config, ip net.IP) *accessResponseURLs {
if ip == nil { if ip == nil {
return nil return nil
+1
View File
@@ -4,6 +4,7 @@ import "github.com/alecthomas/kong"
type CLI struct { type CLI struct {
GenerateSecret GenerateSecret `kong:"cmd,help='Generate new proxy secret'"` GenerateSecret GenerateSecret `kong:"cmd,help='Generate new proxy secret'"`
Doctor Doctor `kong:"cmd,help='Check that proxy can run correctly'"`
Access Access `kong:"cmd,help='Print access information.'"` Access Access `kong:"cmd,help='Print access information.'"`
Run Run `kong:"cmd,help='Run proxy.'"` Run Run `kong:"cmd,help='Run proxy.'"`
SimpleRun SimpleRun `kong:"cmd,help='Run proxy without config file.'"` SimpleRun SimpleRun `kong:"cmd,help='Run proxy without config file.'"`
+381
View File
@@ -0,0 +1,381 @@
package cli
import (
"context"
"errors"
"fmt"
"maps"
"net"
"os"
"slices"
"strconv"
"strings"
"text/template"
"time"
"github.com/9seconds/mtg/v2/essentials"
"github.com/9seconds/mtg/v2/internal/config"
"github.com/9seconds/mtg/v2/internal/utils"
"github.com/9seconds/mtg/v2/mtglib"
"github.com/9seconds/mtg/v2/network/v2"
"github.com/beevik/ntp"
)
var (
tplError = template.Must(
template.New("").Parse(" ‼️ {{ .description }}: {{ .error }}\n"),
)
tplWDeprecatedConfig = template.Must(
template.New("").
Parse(` ⚠️ Option {{ .old | printf "%q" }}{{ if .old_section }} from section [{{ .old_section }}]{{ end }} is deprecated and will be removed in v{{ .when }}. Please use {{ .new | printf "%q" }}{{ if .new_section }} in [{{ .new_section }}] section{{ end }} instead.` + "\n"),
)
tplOTimeSkewness = template.Must(
template.New("").
Parse(" ✅ Time drift is {{ .drift }}, but tolerate-time-skewness is {{ .value }}\n"),
)
tplWTimeSkewness = template.Must(
template.New("").
Parse(" ⚠️ Time drift is {{ .drift }}, but tolerate-time-skewness is {{ .value }}. Please check ntp.\n"),
)
tplETimeSkewness = template.Must(
template.New("").
Parse(" ❌ Time drift is {{ .drift }}, but tolerate-time-skewness is {{ .value }}. You will get many rejected connections!\n"),
)
tplODCConnect = template.Must(
template.New("").Parse(" ✅ DC {{ .dc }}\n"),
)
tplEDCConnect = template.Must(
template.New("").Parse(" ❌ DC {{ .dc }}: {{ .error }}\n"),
)
tplODNSSNIMatch = template.Must(
template.New("").Parse(" ✅ IP address {{ .ip }} matches secret hostname {{ .hostname }}\n"),
)
tplEDNSSNIMatch = template.Must(
template.New("").Parse(" ❌ Hostname {{ .hostname }} {{ if .resolved }}is resolved to {{ .resolved }} addresses, not {{ if .ip4 }}{{ .ip4 }}{{ else }}{{ .ip6 }}{{ end }}{{ else }}cannot be resolved to any host{{ end }}\n"),
)
tplOFrontingDomain = template.Must(
template.New("").Parse(" ✅ {{ .address }} is reachable\n"),
)
tplEFrontingDomain = template.Must(
template.New("").Parse(" ❌ {{ .address }}: {{ .error }}\n"),
)
)
type Doctor struct {
conf *config.Config
ConfigPath string `kong:"arg,required,type='existingfile',help='Path to the configuration file.',name='config-path'"` //nolint: lll
}
func (d *Doctor) Run(cli *CLI, version string) error {
conf, err := utils.ReadConfig(d.ConfigPath)
if err != nil {
return fmt.Errorf("cannot init config: %w", err)
}
d.conf = conf
fmt.Println("Deprecated options")
everythingOK := d.checkDeprecatedConfig()
fmt.Println("Time skewness")
everythingOK = d.checkTimeSkewness() && everythingOK
resolver, err := network.GetDNS(conf.GetDNS())
if err != nil {
return fmt.Errorf("cannot create DNS resolver: %w", err)
}
base := network.New(
resolver,
"",
conf.Network.Timeout.TCP.Get(10*time.Second),
conf.Network.Timeout.HTTP.Get(0),
conf.Network.Timeout.Idle.Get(0),
net.KeepAliveConfig{
Enable: !conf.Network.KeepAlive.Disabled.Get(false),
Idle: conf.Network.KeepAlive.Idle.Get(0),
Interval: conf.Network.KeepAlive.Interval.Get(0),
Count: int(conf.Network.KeepAlive.Count.Get(0)),
},
)
fmt.Println("Validate native network connectivity")
everythingOK = d.checkNetwork(base) && everythingOK
for _, url := range conf.Network.Proxies {
value, err := network.NewProxyNetwork(base, url.Get(nil))
if err != nil {
return err
}
fmt.Printf("Validate network connectivity with proxy %s\n", url.Get(nil))
everythingOK = d.checkNetwork(value) && everythingOK
}
fmt.Println("Validate fronting domain connectivity")
everythingOK = d.checkFrontingDomain(base) && everythingOK
fmt.Println("Validate SNI-DNS match")
everythingOK = d.checkSecretHost(resolver, base) && everythingOK
if !everythingOK {
os.Exit(1)
}
return nil
}
func (d *Doctor) checkDeprecatedConfig() bool {
ok := true
if d.conf.DomainFrontingIP.Value != nil {
ok = false
tplWDeprecatedConfig.Execute(os.Stdout, map[string]string{ //nolint: errcheck
"when": "2.3.0",
"old": "domain-fronting-ip",
"old_section": "",
"new": "ip",
"new_section": "domain-fronting",
})
}
if d.conf.DomainFrontingPort.Value != 0 {
ok = false
tplWDeprecatedConfig.Execute(os.Stdout, map[string]string{ //nolint: errcheck
"when": "2.3.0",
"old": "domain-fronting-port",
"old_section": "",
"new": "port",
"new_section": "domain-fronting",
})
}
if d.conf.DomainFrontingProxyProtocol.Value {
ok = false
tplWDeprecatedConfig.Execute(os.Stdout, map[string]string{ //nolint: errcheck
"when": "2.3.0",
"old": "domain-fronting-proxy-protocol",
"old_section": "",
"new": "proxy-protocol",
"new_section": "domain-fronting",
})
}
if d.conf.Network.DOHIP.Value != nil {
ok = false
tplWDeprecatedConfig.Execute(os.Stdout, map[string]string{ //nolint: errcheck
"when": "2.3.0",
"old": "doh-ip",
"old_section": "network",
"new": "dns",
"new_section": "network",
})
}
if ok {
fmt.Println(" ✅ All good")
}
return ok
}
func (d *Doctor) checkTimeSkewness() bool {
response, err := ntp.Query("0.pool.ntp.org")
if err != nil {
tplError.Execute(os.Stdout, map[string]any{ //nolint: errcheck
"description": "cannot access ntp pool",
"error": err,
})
return false
}
skewness := response.ClockOffset.Abs()
confValue := d.conf.TolerateTimeSkewness.Get(mtglib.DefaultTolerateTimeSkewness)
diff := float64(skewness) / float64(confValue)
tplData := map[string]any{
"drift": response.ClockOffset,
"value": confValue,
}
switch {
case diff < 0.3:
tplOTimeSkewness.Execute(os.Stdout, tplData) //nolint: errcheck
return true
case diff < 0.7:
tplWTimeSkewness.Execute(os.Stdout, tplData) //nolint: errcheck
default:
tplETimeSkewness.Execute(os.Stdout, tplData) //nolint: errcheck
}
return false
}
func (d *Doctor) checkNetwork(ntw mtglib.Network) bool {
dcs := slices.Collect(maps.Keys(essentials.TelegramCoreAddresses))
slices.Sort(dcs)
ok := true
for _, dc := range dcs {
err := d.checkNetworkAddresses(ntw, essentials.TelegramCoreAddresses[dc])
if err == nil {
tplODCConnect.Execute(os.Stdout, map[string]any{ //nolint: errcheck
"dc": dc,
})
} else {
tplEDCConnect.Execute(os.Stdout, map[string]any{ //nolint: errcheck
"dc": dc,
"error": err,
})
ok = false
}
}
return ok
}
func (d *Doctor) checkNetworkAddresses(ntw mtglib.Network, addresses []string) error {
checkAddresses := []string{}
switch d.conf.PreferIP.Get("prefer-ip4") {
case "only-ipv4":
for _, addr := range addresses {
host, _, err := net.SplitHostPort(addr)
if err != nil {
panic(err)
}
if ip := net.ParseIP(host); ip != nil && ip.To4() != nil {
checkAddresses = append(checkAddresses, addr)
}
}
case "only-ipv6":
for _, addr := range addresses {
host, _, err := net.SplitHostPort(addr)
if err != nil {
panic(err)
}
if ip := net.ParseIP(host); ip != nil && ip.To4() == nil {
checkAddresses = append(checkAddresses, addr)
}
}
default:
checkAddresses = addresses
}
if len(checkAddresses) == 0 {
return fmt.Errorf("no suitable addresses after IP version filtering")
}
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
defer cancel()
var (
conn net.Conn
err error
)
for _, addr := range checkAddresses {
conn, err = ntw.DialContext(ctx, "tcp", addr)
if err != nil {
continue
}
conn.Close() //nolint: errcheck
return nil
}
return err
}
func (d *Doctor) checkFrontingDomain(ntw mtglib.Network) bool {
host := d.conf.Secret.Host
if ip := d.conf.GetDomainFrontingIP(nil); ip != "" {
host = ip
}
port := d.conf.GetDomainFrontingPort(mtglib.DefaultDomainFrontingPort)
address := net.JoinHostPort(host, strconv.Itoa(int(port)))
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
defer cancel()
dialer := ntw.NativeDialer()
conn, err := dialer.DialContext(ctx, "tcp", address)
if err != nil {
tplEFrontingDomain.Execute(os.Stdout, map[string]any{ //nolint: errcheck
"address": address,
"error": err,
})
return false
}
conn.Close() //nolint: errcheck
tplOFrontingDomain.Execute(os.Stdout, map[string]any{ //nolint: errcheck
"address": address,
})
return true
}
func (d *Doctor) checkSecretHost(resolver *net.Resolver, ntw mtglib.Network) bool {
addresses, err := resolver.LookupIPAddr(context.Background(), d.conf.Secret.Host)
if err != nil {
tplError.Execute(os.Stdout, map[string]any{ //nolint: errcheck
"description": fmt.Sprintf("cannot resolve DNS name of %s", d.conf.Secret.Host),
"error": err,
})
return false
}
ourIP4 := d.conf.PublicIPv4.Get(nil)
if ourIP4 == nil {
ourIP4 = getIP(ntw, "tcp4")
}
ourIP6 := d.conf.PublicIPv6.Get(nil)
if ourIP6 == nil {
ourIP6 = getIP(ntw, "tcp6")
}
if ourIP4 == nil && ourIP6 == nil {
tplError.Execute(os.Stdout, map[string]any{ //nolint: errcheck
"description": "cannot detect public IP address",
"error": errors.New("cannot detect automatically and public-ipv4/public-ipv6 are not set in config"),
})
return false
}
strAddresses := []string{}
for _, value := range addresses {
if (ourIP4 != nil && value.IP.String() == ourIP4.String()) ||
(ourIP6 != nil && value.IP.String() == ourIP6.String()) {
tplODNSSNIMatch.Execute(os.Stdout, map[string]any{ //nolint: errcheck
"ip": value.IP,
"hostname": d.conf.Secret.Host,
})
return true
}
strAddresses = append(strAddresses, `"`+value.IP.String()+`"`)
}
tplEDNSSNIMatch.Execute(os.Stdout, map[string]any{ //nolint: errcheck
"hostname": d.conf.Secret.Host,
"resolved": strings.Join(strAddresses, ", "),
"ip4": ourIP4,
"ip6": ourIP6,
})
return false
}
+9
View File
@@ -50,6 +50,12 @@ func makeNetwork(conf *config.Config, version string) (mtglib.Network, error) {
conf.Network.Timeout.TCP.Get(0), conf.Network.Timeout.TCP.Get(0),
conf.Network.Timeout.HTTP.Get(0), conf.Network.Timeout.HTTP.Get(0),
conf.Network.Timeout.Idle.Get(0), conf.Network.Timeout.Idle.Get(0),
net.KeepAliveConfig{
Enable: !conf.Network.KeepAlive.Disabled.Get(false),
Idle: conf.Network.KeepAlive.Idle.Get(0),
Interval: conf.Network.KeepAlive.Interval.Get(0),
Count: int(conf.Network.KeepAlive.Count.Get(0)),
},
) )
proxyDialers := make([]mtglib.Network, len(conf.Network.Proxies)) proxyDialers := make([]mtglib.Network, len(conf.Network.Proxies))
@@ -253,6 +259,7 @@ func runProxy(conf *config.Config, version string) error { //nolint: funlen
EventStream: eventStream, EventStream: eventStream,
Secret: conf.Secret, Secret: conf.Secret,
Concurrency: conf.GetConcurrency(mtglib.DefaultConcurrency),
DomainFrontingPort: conf.GetDomainFrontingPort(mtglib.DefaultDomainFrontingPort), DomainFrontingPort: conf.GetDomainFrontingPort(mtglib.DefaultDomainFrontingPort),
DomainFrontingIP: conf.GetDomainFrontingIP(nil), DomainFrontingIP: conf.GetDomainFrontingIP(nil),
DomainFrontingProxyProtocol: conf.GetDomainFrontingProxyProtocol(false), DomainFrontingProxyProtocol: conf.GetDomainFrontingProxyProtocol(false),
@@ -261,6 +268,8 @@ func runProxy(conf *config.Config, version string) error { //nolint: funlen
AllowFallbackOnUnknownDC: conf.AllowFallbackOnUnknownDC.Get(false), AllowFallbackOnUnknownDC: conf.AllowFallbackOnUnknownDC.Get(false),
TolerateTimeSkewness: conf.TolerateTimeSkewness.Value, TolerateTimeSkewness: conf.TolerateTimeSkewness.Value,
IdleTimeout: conf.Network.Timeout.Idle.Get(mtglib.DefaultIdleTimeout),
HandshakeTimeout: conf.Network.Timeout.Handshake.Get(mtglib.DefaultHandshakeTimeout),
DoppelGangerURLs: doppelGangerURLs, DoppelGangerURLs: doppelGangerURLs,
DoppelGangerPerRaid: conf.Defense.Doppelganger.Repeats.Get(mtglib.DoppelGangerPerRaid), DoppelGangerPerRaid: conf.Defense.Doppelganger.Repeats.Get(mtglib.DoppelGangerPerRaid),
+51
View File
@@ -0,0 +1,51 @@
package cli
import (
"context"
"io"
"net"
"net/http"
"strings"
"github.com/9seconds/mtg/v2/essentials"
"github.com/9seconds/mtg/v2/mtglib"
)
func getIP(ntw mtglib.Network, protocol string) net.IP {
dialer := ntw.NativeDialer()
client := ntw.MakeHTTPClient(func(ctx context.Context, network, address string) (essentials.Conn, error) {
conn, err := dialer.DialContext(ctx, protocol, address)
if err != nil {
return nil, err
}
return essentials.WrapNetConn(conn), err
})
req, err := http.NewRequest(http.MethodGet, "https://ifconfig.co", nil) //nolint: noctx
if err != nil {
panic(err)
}
req.Header.Add("Accept", "text/plain")
resp, err := client.Do(req)
if err != nil {
return nil
}
if resp.StatusCode != http.StatusOK {
return nil
}
defer func() {
io.Copy(io.Discard, resp.Body) //nolint: errcheck
resp.Body.Close() //nolint: errcheck
}()
data, err := io.ReadAll(resp.Body)
if err != nil {
return nil
}
return net.ParseIP(strings.TrimSpace(string(data)))
}
+16
View File
@@ -35,6 +35,8 @@ type Config struct {
DomainFrontingProxyProtocol TypeBool `json:"domainFrontingProxyProtocol"` DomainFrontingProxyProtocol TypeBool `json:"domainFrontingProxyProtocol"`
TolerateTimeSkewness TypeDuration `json:"tolerateTimeSkewness"` TolerateTimeSkewness TypeDuration `json:"tolerateTimeSkewness"`
Concurrency TypeConcurrency `json:"concurrency"` Concurrency TypeConcurrency `json:"concurrency"`
PublicIPv4 TypeIP `json:"publicIpv4"`
PublicIPv6 TypeIP `json:"publicIpv6"`
DomainFronting struct { DomainFronting struct {
IP TypeIP `json:"ip"` IP TypeIP `json:"ip"`
Port TypePort `json:"port"` Port TypePort `json:"port"`
@@ -61,7 +63,14 @@ type Config struct {
TCP TypeDuration `json:"tcp"` TCP TypeDuration `json:"tcp"`
HTTP TypeDuration `json:"http"` HTTP TypeDuration `json:"http"`
Idle TypeDuration `json:"idle"` Idle TypeDuration `json:"idle"`
Handshake TypeDuration `json:"handshake"`
} `json:"timeout"` } `json:"timeout"`
KeepAlive struct {
Disabled TypeBool `json:"disabled"`
Idle TypeDuration `json:"idle"`
Interval TypeDuration `json:"interval"`
Count TypeConcurrency `json:"count"`
} `json:"keepAlive"`
DOHIP TypeIP `json:"dohIp"` DOHIP TypeIP `json:"dohIp"`
DNS TypeDNSURI `json:"dns"` DNS TypeDNSURI `json:"dns"`
Proxies []TypeProxyURL `json:"proxies"` Proxies []TypeProxyURL `json:"proxies"`
@@ -84,6 +93,13 @@ type Config struct {
} `json:"stats"` } `json:"stats"`
} }
func (c *Config) GetConcurrency(defaultValue uint) uint {
if concurrency := c.Concurrency.Get(0); concurrency != 0 {
return concurrency
}
return c.Concurrency.Get(defaultValue)
}
func (c *Config) GetDNS() *url.URL { func (c *Config) GetDNS() *url.URL {
var dohURL *url.URL var dohURL *url.URL
+26
View File
@@ -42,6 +42,32 @@ func (suite *ConfigTestSuite) TestParseMinimalConfig() {
suite.Equal("0.0.0.0:3128", conf.BindTo.String()) suite.Equal("0.0.0.0:3128", conf.BindTo.String())
} }
func (suite *ConfigTestSuite) TestParsePublicIP() {
conf, err := config.Parse(suite.ReadConfig("public_ip.toml"))
suite.NoError(err)
suite.Equal("203.0.113.1", conf.PublicIPv4.Get(nil).String())
suite.Equal("2001:db8::1", conf.PublicIPv6.Get(nil).String())
}
func (suite *ConfigTestSuite) TestParsePublicIPv4Only() {
conf, err := config.Parse(suite.ReadConfig("public_ip_v4_only.toml"))
suite.NoError(err)
suite.Equal("203.0.113.1", conf.PublicIPv4.Get(nil).String())
suite.Nil(conf.PublicIPv6.Get(nil))
}
func (suite *ConfigTestSuite) TestParsePublicIPInvalid() {
_, err := config.Parse(suite.ReadConfig("public_ip_invalid.toml"))
suite.Error(err)
}
func (suite *ConfigTestSuite) TestParsePublicIPNotSet() {
conf, err := config.Parse(suite.ReadConfig("minimal.toml"))
suite.NoError(err)
suite.Nil(conf.PublicIPv4.Get(nil))
suite.Nil(conf.PublicIPv6.Get(nil))
}
func (suite *ConfigTestSuite) TestString() { func (suite *ConfigTestSuite) TestString() {
conf, err := config.Parse(suite.ReadConfig("minimal.toml")) conf, err := config.Parse(suite.ReadConfig("minimal.toml"))
suite.NoError(err) suite.NoError(err)
+9
View File
@@ -21,6 +21,8 @@ type tomlConfig struct {
DomainFrontingProxyProtocol bool `toml:"domain-fronting-proxy-protocol" json:"domainFrontingProxyProtocol,omitempty"` DomainFrontingProxyProtocol bool `toml:"domain-fronting-proxy-protocol" json:"domainFrontingProxyProtocol,omitempty"`
TolerateTimeSkewness string `toml:"tolerate-time-skewness" json:"tolerateTimeSkewness,omitempty"` TolerateTimeSkewness string `toml:"tolerate-time-skewness" json:"tolerateTimeSkewness,omitempty"`
Concurrency uint `toml:"concurrency" json:"concurrency,omitempty"` Concurrency uint `toml:"concurrency" json:"concurrency,omitempty"`
PublicIPv4 string `toml:"public-ipv4" json:"publicIpv4,omitempty"`
PublicIPv6 string `toml:"public-ipv6" json:"publicIpv6,omitempty"`
DomainFronting struct { DomainFronting struct {
IP string `toml:"ip" json:"ip,omitempty"` IP string `toml:"ip" json:"ip,omitempty"`
Port uint `toml:"port" json:"port,omitempty"` Port uint `toml:"port" json:"port,omitempty"`
@@ -56,7 +58,14 @@ type tomlConfig struct {
TCP string `toml:"tcp" json:"tcp,omitempty"` TCP string `toml:"tcp" json:"tcp,omitempty"`
HTTP string `toml:"http" json:"http,omitempty"` HTTP string `toml:"http" json:"http,omitempty"`
Idle string `toml:"idle" json:"idle,omitempty"` Idle string `toml:"idle" json:"idle,omitempty"`
Handshake string `toml:"handshake" json:"handshake,omitempty"`
} `toml:"timeout" json:"timeout,omitempty"` } `toml:"timeout" json:"timeout,omitempty"`
KeepAlive struct {
Disabled bool `toml:"disabled" json:"disabled,omitempty"`
Idle string `toml:"idle" json:"idle,omitempty"`
Interval string `toml:"interval" json:"interval,omitempty"`
Count uint `toml:"count" json:"count,omitempty"`
} `toml:"keep-alive" json:"keepAlive,omitempty"`
DOHIP string `toml:"doh-ip" json:"dohIp,omitempty"` DOHIP string `toml:"doh-ip" json:"dohIp,omitempty"`
DNS string `toml:"dns" json:"dns,omitempty"` DNS string `toml:"dns" json:"dns,omitempty"`
Proxies []string `toml:"proxies" json:"proxies,omitempty"` Proxies []string `toml:"proxies" json:"proxies,omitempty"`
+4
View File
@@ -0,0 +1,4 @@
secret = "7oe1GqLy6TBc38CV3jx7q09nb29nbGUuY29t"
bind-to = "0.0.0.0:3128"
public-ipv4 = "203.0.113.1"
public-ipv6 = "2001:db8::1"
+3
View File
@@ -0,0 +1,3 @@
secret = "7oe1GqLy6TBc38CV3jx7q09nb29nbGUuY29t"
bind-to = "0.0.0.0:3128"
public-ipv4 = "not-an-ip"
+3
View File
@@ -0,0 +1,3 @@
secret = "7oe1GqLy6TBc38CV3jx7q09nb29nbGUuY29t"
bind-to = "0.0.0.0:3128"
public-ipv4 = "203.0.113.1"
+17
View File
@@ -14,6 +14,23 @@ import (
) )
func main() { func main() {
// this runs profiling server. To enable it, build with prof tag
// $ go build -tags prof
//
// Then you can pass a port using MTG_PROF_PORT environment variable.
// Default is 6000
// $ MTG_PROF_PORT=6000 mtg run config.toml
//
// It will run a webserver with profiling data on
// localhost:${MTG_PROF_PORT:-6000}.
//
// To collect PGO do following:
// $ curl -o default.pgo 'http://localhost:6000/debug/pprof/profile?seconds=300'
//
// See also https://pkg.go.dev/net/http/pprof
// https://go.dev/blog/pprof
runProfile()
cli := &cli.CLI{} cli := &cli.CLI{}
ctx := kong.Parse(cli, kong.Vars{ ctx := kong.Parse(cli, kong.Vars{
"version": getVersion(), "version": getVersion(),
+96 -17
View File
@@ -1,11 +1,36 @@
# @generated - this file is auto-generated by `mise lock` https://mise.jdx.dev/dev-tools/mise-lock.html
[[tools.go]] [[tools.go]]
version = "1.26.1" version = "1.26.1"
backend = "core:go" backend = "core:go"
"platforms.linux-arm64" = { checksum = "sha256:a290581cfe4fe28ddd737dde3095f3dbeb7f2e4065cab4eae44dfc53b760c2f7", url = "https://dl.google.com/go/go1.26.1.linux-arm64.tar.gz"}
"platforms.linux-x64" = { checksum = "sha256:031f088e5d955bab8657ede27ad4e3bc5b7c1ba281f05f245bcc304f327c987a", url = "https://dl.google.com/go/go1.26.1.linux-amd64.tar.gz"} [tools.go."platforms.linux-arm64"]
"platforms.macos-arm64" = { checksum = "sha256:353df43a7811ce284c8938b5f3c7df40b7bfb6f56cb165b150bc40b5e2dd541f", url = "https://dl.google.com/go/go1.26.1.darwin-arm64.tar.gz"} checksum = "sha256:a290581cfe4fe28ddd737dde3095f3dbeb7f2e4065cab4eae44dfc53b760c2f7"
"platforms.macos-x64" = { checksum = "sha256:65773dab2f8cc4cd23d93ba6d0a805de150ca0b78378879292be0b903b8cdd08", url = "https://dl.google.com/go/go1.26.1.darwin-amd64.tar.gz"} url = "https://dl.google.com/go/go1.26.1.linux-arm64.tar.gz"
"platforms.windows-x64" = { checksum = "sha256:9b68112c913f45b7aebbf13c036721264bbba7e03a642f8f7490c561eebd1ecc", url = "https://dl.google.com/go/go1.26.1.windows-amd64.zip"}
[tools.go."platforms.linux-arm64-musl"]
checksum = "sha256:a290581cfe4fe28ddd737dde3095f3dbeb7f2e4065cab4eae44dfc53b760c2f7"
url = "https://dl.google.com/go/go1.26.1.linux-arm64.tar.gz"
[tools.go."platforms.linux-x64"]
checksum = "sha256:031f088e5d955bab8657ede27ad4e3bc5b7c1ba281f05f245bcc304f327c987a"
url = "https://dl.google.com/go/go1.26.1.linux-amd64.tar.gz"
[tools.go."platforms.linux-x64-musl"]
checksum = "sha256:031f088e5d955bab8657ede27ad4e3bc5b7c1ba281f05f245bcc304f327c987a"
url = "https://dl.google.com/go/go1.26.1.linux-amd64.tar.gz"
[tools.go."platforms.macos-arm64"]
checksum = "sha256:353df43a7811ce284c8938b5f3c7df40b7bfb6f56cb165b150bc40b5e2dd541f"
url = "https://dl.google.com/go/go1.26.1.darwin-arm64.tar.gz"
[tools.go."platforms.macos-x64"]
checksum = "sha256:65773dab2f8cc4cd23d93ba6d0a805de150ca0b78378879292be0b903b8cdd08"
url = "https://dl.google.com/go/go1.26.1.darwin-amd64.tar.gz"
[tools.go."platforms.windows-x64"]
checksum = "sha256:9b68112c913f45b7aebbf13c036721264bbba7e03a642f8f7490c561eebd1ecc"
url = "https://dl.google.com/go/go1.26.1.windows-amd64.zip"
[[tools."go:golang.org/x/pkgsite/cmd/pkgsite"]] [[tools."go:golang.org/x/pkgsite/cmd/pkgsite"]]
version = "latest" version = "latest"
@@ -24,19 +49,73 @@ version = "0.9.2"
backend = "go:mvdan.cc/gofumpt" backend = "go:mvdan.cc/gofumpt"
[[tools.golangci-lint]] [[tools.golangci-lint]]
version = "2.11.3" version = "2.11.4"
backend = "aqua:golangci/golangci-lint" backend = "aqua:golangci/golangci-lint"
"platforms.linux-arm64" = { checksum = "sha256:ee3d95f301359e7d578e6d99c8ad5aeadbabc5a13009a30b2b0df11c8058afe9", url = "https://github.com/golangci/golangci-lint/releases/download/v2.11.3/golangci-lint-2.11.3-linux-arm64.tar.gz"}
"platforms.linux-x64" = { checksum = "sha256:87bb8cddbcc825d5778b64e8a91b46c0526b247f4e2f2904dea74ec7450475d1", url = "https://github.com/golangci/golangci-lint/releases/download/v2.11.3/golangci-lint-2.11.3-linux-amd64.tar.gz"} [tools.golangci-lint."platforms.linux-arm64"]
"platforms.macos-arm64" = { checksum = "sha256:30ee39979c516b9d1adca289a3f93429d130c4c0fda5e57d637850894221f6cc", url = "https://github.com/golangci/golangci-lint/releases/download/v2.11.3/golangci-lint-2.11.3-darwin-arm64.tar.gz"} checksum = "sha256:3bcfa2e6f3d32b2bf5cd75eaa876447507025e0303698633f722a05331988db4"
"platforms.macos-x64" = { checksum = "sha256:f93bda1f2cc981fd1326464020494be62f387bbf262706e1b3b644e5afacc440", url = "https://github.com/golangci/golangci-lint/releases/download/v2.11.3/golangci-lint-2.11.3-darwin-amd64.tar.gz"} url = "https://github.com/golangci/golangci-lint/releases/download/v2.11.4/golangci-lint-2.11.4-linux-arm64.tar.gz"
"platforms.windows-x64" = { checksum = "sha256:cd42e890176bc5cfeb36225a77e66b9410ddd3a59a03551e23f6b210d29e1f67", url = "https://github.com/golangci/golangci-lint/releases/download/v2.11.3/golangci-lint-2.11.3-windows-amd64.zip"}
[tools.golangci-lint."platforms.linux-arm64-musl"]
checksum = "sha256:3bcfa2e6f3d32b2bf5cd75eaa876447507025e0303698633f722a05331988db4"
url = "https://github.com/golangci/golangci-lint/releases/download/v2.11.4/golangci-lint-2.11.4-linux-arm64.tar.gz"
[tools.golangci-lint."platforms.linux-x64"]
checksum = "sha256:200c5b7503f67b59a6743ccf32133026c174e272b930ee79aa2aa6f37aca7ef1"
url = "https://github.com/golangci/golangci-lint/releases/download/v2.11.4/golangci-lint-2.11.4-linux-amd64.tar.gz"
[tools.golangci-lint."platforms.linux-x64-musl"]
checksum = "sha256:200c5b7503f67b59a6743ccf32133026c174e272b930ee79aa2aa6f37aca7ef1"
url = "https://github.com/golangci/golangci-lint/releases/download/v2.11.4/golangci-lint-2.11.4-linux-amd64.tar.gz"
[tools.golangci-lint."platforms.macos-arm64"]
checksum = "sha256:02db2a2dae8b26812e53b0688a6f617e3ef1f489790e829ea22862cf76945675"
url = "https://github.com/golangci/golangci-lint/releases/download/v2.11.4/golangci-lint-2.11.4-darwin-arm64.tar.gz"
provenance = "github-attestations"
[tools.golangci-lint."platforms.macos-x64"]
checksum = "sha256:c900d4048db75d1edfd550fd11cf6a9b3008e7caa8e119fcddbc700412d63e60"
url = "https://github.com/golangci/golangci-lint/releases/download/v2.11.4/golangci-lint-2.11.4-darwin-amd64.tar.gz"
[tools.golangci-lint."platforms.windows-x64"]
checksum = "sha256:4932cfca5e75bf60fe1c576edf459e5e809e6644664a068185d64b84af3fad9e"
url = "https://github.com/golangci/golangci-lint/releases/download/v2.11.4/golangci-lint-2.11.4-windows-amd64.zip"
[[tools.goreleaser]] [[tools.goreleaser]]
version = "2.14.3" version = "2.15.2"
backend = "aqua:goreleaser/goreleaser" backend = "aqua:goreleaser/goreleaser"
"platforms.linux-arm64" = { checksum = "sha256:581a10e53c1176b3e81ee45cf531e02dbf899db0bc7b795669347df4276ce948", url = "https://github.com/goreleaser/goreleaser/releases/download/v2.14.3/goreleaser_Linux_arm64.tar.gz"}
"platforms.linux-x64" = { checksum = "sha256:dc7faeeeb6da8bdfda788626263a4ae725892a8c7504b975c3234127d4a44579", url = "https://github.com/goreleaser/goreleaser/releases/download/v2.14.3/goreleaser_Linux_x86_64.tar.gz"} [tools.goreleaser."platforms.linux-arm64"]
"platforms.macos-arm64" = { checksum = "sha256:3507798489e107a78aff36b169de48148a335ac26eb3161608d905f3f3a957bd", url = "https://github.com/goreleaser/goreleaser/releases/download/v2.14.3/goreleaser_Darwin_all.tar.gz"} checksum = "sha256:5db66761a98f6693161e49e1a95d28d2673a892ba60cb4a5e16736cafd41c4c9"
"platforms.macos-x64" = { checksum = "sha256:3507798489e107a78aff36b169de48148a335ac26eb3161608d905f3f3a957bd", url = "https://github.com/goreleaser/goreleaser/releases/download/v2.14.3/goreleaser_Darwin_all.tar.gz"} url = "https://github.com/goreleaser/goreleaser/releases/download/v2.15.2/goreleaser_Linux_arm64.tar.gz"
"platforms.windows-x64" = { checksum = "sha256:3deea8ff471aa258a2d99f3e5302971d7028647ae8ddaf103257a8113e485a31", url = "https://github.com/goreleaser/goreleaser/releases/download/v2.14.3/goreleaser_Windows_x86_64.zip"} provenance = "cosign"
[tools.goreleaser."platforms.linux-arm64-musl"]
checksum = "sha256:5db66761a98f6693161e49e1a95d28d2673a892ba60cb4a5e16736cafd41c4c9"
url = "https://github.com/goreleaser/goreleaser/releases/download/v2.15.2/goreleaser_Linux_arm64.tar.gz"
provenance = "cosign"
[tools.goreleaser."platforms.linux-x64"]
checksum = "sha256:0ebdbf0353aba566b969dde746cc4e4806f96c27aa2f3971b229a9df7611fedc"
url = "https://github.com/goreleaser/goreleaser/releases/download/v2.15.2/goreleaser_Linux_x86_64.tar.gz"
provenance = "cosign"
[tools.goreleaser."platforms.linux-x64-musl"]
checksum = "sha256:0ebdbf0353aba566b969dde746cc4e4806f96c27aa2f3971b229a9df7611fedc"
url = "https://github.com/goreleaser/goreleaser/releases/download/v2.15.2/goreleaser_Linux_x86_64.tar.gz"
provenance = "cosign"
[tools.goreleaser."platforms.macos-arm64"]
checksum = "sha256:0e6bd67688ac949780bf1166813a91f89856898ef4c40d7d46c2c74ebaa4b9ee"
url = "https://github.com/goreleaser/goreleaser/releases/download/v2.15.2/goreleaser_Darwin_all.tar.gz"
provenance = "github-attestations"
[tools.goreleaser."platforms.macos-x64"]
checksum = "sha256:0e6bd67688ac949780bf1166813a91f89856898ef4c40d7d46c2c74ebaa4b9ee"
url = "https://github.com/goreleaser/goreleaser/releases/download/v2.15.2/goreleaser_Darwin_all.tar.gz"
provenance = "cosign"
[tools.goreleaser."platforms.windows-x64"]
checksum = "sha256:7459832946dbe122c144f8d7f87484d8572ca005b779310aa6bb03346e8de17a"
url = "https://github.com/goreleaser/goreleaser/releases/download/v2.15.2/goreleaser_Windows_x86_64.zip"
provenance = "cosign"
+64
View File
@@ -3,9 +3,12 @@ package mtglib
import ( import (
"bytes" "bytes"
"context" "context"
"errors"
"fmt" "fmt"
"io" "io"
"net" "net"
"sync/atomic"
"time"
"github.com/9seconds/mtg/v2/essentials" "github.com/9seconds/mtg/v2/essentials"
"github.com/pires/go-proxyproto" "github.com/pires/go-proxyproto"
@@ -95,3 +98,64 @@ func newConnProxyProtocol(source, target essentials.Conn) *connProxyProtocol {
sourceAddr: source.RemoteAddr(), sourceAddr: source.RemoteAddr(),
} }
} }
// idleTracker is a shared idle tracker for a pair of relay connections.
// Both directions update the same timestamp so that activity in one direction
// prevents the other (idle) direction from timing out.
type idleTracker struct {
lastActive atomic.Pointer[time.Time]
timeout time.Duration
}
func newIdleTracker(timeout time.Duration) *idleTracker {
t := &idleTracker{timeout: timeout}
t.touch()
return t
}
func (t *idleTracker) touch() {
stamp := time.Now()
t.lastActive.Store(&stamp)
}
func (t *idleTracker) isIdle() bool {
return time.Since(*t.lastActive.Load()) >= t.timeout
}
type connIdleTimeout struct {
essentials.Conn
tracker *idleTracker
}
func (c connIdleTimeout) Read(b []byte) (int, error) {
var netErr net.Error
for {
c.SetReadDeadline(time.Now().Add(c.tracker.timeout)) //nolint: errcheck
n, err := c.Conn.Read(b)
switch {
case err == nil:
c.tracker.touch()
return n, nil
case errors.As(err, &netErr) && netErr.Timeout() && !c.tracker.isIdle():
continue
}
return n, err
}
}
func (c connIdleTimeout) Write(b []byte) (int, error) {
c.SetWriteDeadline(time.Now().Add(c.tracker.timeout)) //nolint: errcheck
n, err := c.Conn.Write(b)
if n > 0 {
c.tracker.touch()
}
return n, err //nolint: wrapcheck
}
+151
View File
@@ -16,6 +16,12 @@ import (
"github.com/stretchr/testify/suite" "github.com/stretchr/testify/suite"
) )
type netTimeoutError struct{}
func (e netTimeoutError) Error() string { return "i/o timeout" }
func (e netTimeoutError) Timeout() bool { return true }
func (e netTimeoutError) Temporary() bool { return true }
type ConnRewindBaseConn struct { type ConnRewindBaseConn struct {
testlib.EssentialsConnMock testlib.EssentialsConnMock
@@ -291,6 +297,141 @@ func (suite *ConnProxyProtocolTestSuite) TearDownTest() {
suite.targetConnMock.AssertExpectations(suite.T()) suite.targetConnMock.AssertExpectations(suite.T())
} }
type IdleTrackerTestSuite struct {
suite.Suite
}
func (suite *IdleTrackerTestSuite) TestNewNotIdle() {
tracker := newIdleTracker(time.Second)
suite.False(tracker.isIdle())
}
func (suite *IdleTrackerTestSuite) TestIdleAfterTimeout() {
tracker := newIdleTracker(10 * time.Millisecond)
time.Sleep(20 * time.Millisecond)
suite.True(tracker.isIdle())
}
func (suite *IdleTrackerTestSuite) TestTouchResetsIdle() {
tracker := newIdleTracker(50 * time.Millisecond)
time.Sleep(30 * time.Millisecond)
tracker.touch()
suite.False(tracker.isIdle())
}
type ConnIdleTimeoutTestSuite struct {
suite.Suite
connMock *testlib.EssentialsConnMock
tracker *idleTracker
conn connIdleTimeout
}
func (suite *ConnIdleTimeoutTestSuite) SetupTest() {
suite.connMock = &testlib.EssentialsConnMock{}
suite.tracker = newIdleTracker(time.Second)
suite.conn = connIdleTimeout{
Conn: suite.connMock,
tracker: suite.tracker,
}
}
func (suite *ConnIdleTimeoutTestSuite) TearDownTest() {
suite.connMock.AssertExpectations(suite.T())
}
func (suite *ConnIdleTimeoutTestSuite) TestReadOk() {
suite.connMock.On("SetReadDeadline", mock.Anything).Return(nil)
suite.connMock.On("Read", mock.Anything).Once().Return(5, nil)
n, err := suite.conn.Read(make([]byte, 10))
suite.NoError(err)
suite.Equal(5, n)
}
func (suite *ConnIdleTimeoutTestSuite) TestReadNonTimeoutErr() {
suite.connMock.On("SetReadDeadline", mock.Anything).Return(nil)
suite.connMock.On("Read", mock.Anything).Once().Return(0, io.EOF)
n, err := suite.conn.Read(make([]byte, 10))
suite.True(errors.Is(err, io.EOF))
suite.Equal(0, n)
}
func (suite *ConnIdleTimeoutTestSuite) TestReadTimeoutRetriesWhenNotIdle() {
suite.connMock.On("SetReadDeadline", mock.Anything).Return(nil)
suite.connMock.On("Read", mock.Anything).Once().Return(0, netTimeoutError{})
suite.connMock.On("Read", mock.Anything).Once().Return(5, nil)
n, err := suite.conn.Read(make([]byte, 10))
suite.NoError(err)
suite.Equal(5, n)
}
func (suite *ConnIdleTimeoutTestSuite) TestReadTimeoutClosesWhenIdle() {
suite.tracker = newIdleTracker(time.Millisecond)
suite.conn = connIdleTimeout{
Conn: suite.connMock,
tracker: suite.tracker,
}
time.Sleep(5 * time.Millisecond)
suite.connMock.On("SetReadDeadline", mock.Anything).Return(nil)
suite.connMock.On("Read", mock.Anything).Once().Return(0, netTimeoutError{})
n, err := suite.conn.Read(make([]byte, 10))
suite.Equal(0, n)
netErr, ok := err.(net.Error) //nolint: errorlint
suite.True(ok)
suite.True(netErr.Timeout())
}
func (suite *ConnIdleTimeoutTestSuite) TestSharedTrackerPreventsFalseTimeout() {
connMock2 := &testlib.EssentialsConnMock{}
conn2 := connIdleTimeout{
Conn: connMock2,
tracker: suite.tracker,
}
connMock2.On("SetWriteDeadline", mock.Anything).Return(nil)
connMock2.On("Write", mock.Anything).Once().Return(5, nil)
_, _ = conn2.Write(make([]byte, 5))
suite.connMock.On("SetReadDeadline", mock.Anything).Return(nil)
suite.connMock.On("Read", mock.Anything).Once().Return(0, netTimeoutError{})
suite.connMock.On("Read", mock.Anything).Once().Return(3, nil)
n, err := suite.conn.Read(make([]byte, 10))
suite.NoError(err)
suite.Equal(3, n)
connMock2.AssertExpectations(suite.T())
}
func (suite *ConnIdleTimeoutTestSuite) TestWriteOk() {
suite.connMock.On("SetWriteDeadline", mock.Anything).Return(nil)
suite.connMock.On("Write", mock.Anything).Once().Return(5, nil)
n, err := suite.conn.Write(make([]byte, 5))
suite.NoError(err)
suite.Equal(5, n)
}
func (suite *ConnIdleTimeoutTestSuite) TestWriteErr() {
suite.connMock.On("SetWriteDeadline", mock.Anything).Return(nil)
suite.connMock.On("Write", mock.Anything).Once().Return(0, io.EOF)
n, err := suite.conn.Write(make([]byte, 5))
suite.True(errors.Is(err, io.EOF))
suite.Equal(0, n)
}
func TestConnTraffic(t *testing.T) { func TestConnTraffic(t *testing.T) {
t.Parallel() t.Parallel()
suite.Run(t, &ConnTrafficTestSuite{}) suite.Run(t, &ConnTrafficTestSuite{})
@@ -305,3 +446,13 @@ func TestConnProxyProtocol(t *testing.T) {
t.Parallel() t.Parallel()
suite.Run(t, &ConnProxyProtocolTestSuite{}) suite.Run(t, &ConnProxyProtocolTestSuite{})
} }
func TestIdleTracker(t *testing.T) {
t.Parallel()
suite.Run(t, &IdleTrackerTestSuite{})
}
func TestConnIdleTimeout(t *testing.T) {
t.Parallel()
suite.Run(t, &ConnIdleTimeoutTestSuite{})
}
+7 -2
View File
@@ -77,8 +77,13 @@ const (
// DefaultIdleTimeout is a default timeout for closing a connection in case of // DefaultIdleTimeout is a default timeout for closing a connection in case of
// idling. // idling.
// //
// Deprecated: no longer in use because of changed TCP relay algorithm. // Set to 5 minutes to survive typical mobile sleep periods (2-5 min) and
DefaultIdleTimeout = time.Minute // avoid racing with MTProto ping_delay_disconnect (~60s interval).
DefaultIdleTimeout = 5 * time.Minute
// DefaultHandshakeTimeout defines a time period during which the
// all handshake ceremonies must be completed.
DefaultHandshakeTimeout = 10 * time.Second
// DefaultTolerateTimeSkewness is a default timeout for time skewness on a // DefaultTolerateTimeSkewness is a default timeout for time skewness on a
// faketls timeout verification. // faketls timeout verification.
+35 -42
View File
@@ -2,7 +2,10 @@ package dc
import ( import (
"context" "context"
"net"
"time" "time"
"github.com/9seconds/mtg/v2/essentials"
) )
type preferIP uint8 type preferIP uint8
@@ -39,46 +42,36 @@ type Updater interface {
} }
// https://github.com/telegramdesktop/tdesktop/blob/master/Telegram/SourceFiles/mtproto/mtproto_dc_options.cpp#L30 // https://github.com/telegramdesktop/tdesktop/blob/master/Telegram/SourceFiles/mtproto/mtproto_dc_options.cpp#L30
var defaultDCAddrSet = dcAddrSet{ var defaultDCAddrSet = (func() dcAddrSet {
v4: map[int][]Addr{ addrSet := dcAddrSet{
1: { v4: make(map[int][]Addr),
{Network: "tcp4", Address: "149.154.175.50:443"}, v6: make(map[int][]Addr),
},
2: {
{Network: "tcp4", Address: "149.154.167.51:443"},
{Network: "tcp4", Address: "95.161.76.100:443"},
},
3: {
{Network: "tcp4", Address: "149.154.175.100:443"},
},
4: {
{Network: "tcp4", Address: "149.154.167.91:443"},
},
5: {
{Network: "tcp4", Address: "149.154.171.5:443"},
},
203: {
{Network: "tcp4", Address: "91.105.192.100:443"},
},
},
v6: map[int][]Addr{
1: {
{Network: "tcp6", Address: "[2001:b28:f23d:f001::a]:443"},
},
2: {
{Network: "tcp6", Address: "[2001:67c:04e8:f002::a]:443"},
},
3: {
{Network: "tcp6", Address: "[2001:b28:f23d:f003::a]:443"},
},
4: {
{Network: "tcp6", Address: "[2001:67c:04e8:f004::a]:443"},
},
5: {
{Network: "tcp6", Address: "[2001:b28:f23f:f005::a]:443"},
},
203: {
{Network: "tcp6", Address: "[2a0a:f280:0203:000a:5000:0000:0000:0100]:443"},
},
},
} }
for dcid, ips := range essentials.TelegramCoreAddresses {
for _, addr := range ips {
host, _, err := net.SplitHostPort(addr)
if err != nil {
panic(err)
}
ip := net.ParseIP(host)
if ip == nil {
panic(addr)
}
if ip.To4() == nil {
addrSet.v6[dcid] = append(addrSet.v6[dcid], Addr{
Network: "tcp6",
Address: addr,
})
} else {
addrSet.v4[dcid] = append(addrSet.v4[dcid], Addr{
Network: "tcp4",
Address: addr,
})
}
}
}
return addrSet
})()
+1 -1
View File
@@ -20,7 +20,7 @@ func (t *Telegram) GetAddresses(dc int) []Addr {
case preferIPOnlyIPv4: case preferIPOnlyIPv4:
return t.view.getV4(dc) return t.view.getV4(dc)
case preferIPOnlyIPv6: case preferIPOnlyIPv6:
return t.view.getV4(dc) return t.view.getV6(dc)
case preferIPPreferIPv4: case preferIPPreferIPv4:
return append(t.view.getV4(dc), t.view.getV6(dc)...) return append(t.view.getV4(dc), t.view.getV6(dc)...)
} }
+6 -2
View File
@@ -5,15 +5,19 @@ type dcView struct {
} }
func (d dcView) getV4(dc int) []Addr { func (d dcView) getV4(dc int) []Addr {
addrs := d.publicConfigs.getV4(dc) var addrs []Addr
addrs = append(addrs, defaultDCAddrSet.getV4(dc)...) addrs = append(addrs, defaultDCAddrSet.getV4(dc)...)
addrs = append(addrs, d.publicConfigs.getV4(dc)...)
return addrs return addrs
} }
func (d dcView) getV6(dc int) []Addr { func (d dcView) getV6(dc int) []Addr {
addrs := d.publicConfigs.getV6(dc) var addrs []Addr
addrs = append(addrs, defaultDCAddrSet.getV6(dc)...) addrs = append(addrs, defaultDCAddrSet.getV6(dc)...)
addrs = append(addrs, d.publicConfigs.getV6(dc)...)
return addrs return addrs
} }
-35
View File
@@ -1,35 +0,0 @@
package doppel
import (
"context"
"time"
)
type Clock struct {
stats *Stats
tick chan struct{}
}
func (c Clock) Start(ctx context.Context) {
tickTock := time.NewTimer(c.stats.Delay())
defer func() {
tickTock.Stop()
select {
case <-tickTock.C:
default:
}
}()
for {
select {
case <-ctx.Done():
return
case <-tickTock.C:
select {
case <-ctx.Done():
case c.tick <- struct{}{}:
}
tickTock.Reset(c.stats.Delay())
}
}
}
-80
View File
@@ -1,80 +0,0 @@
package doppel
import (
"context"
"sync"
"testing"
"time"
"github.com/stretchr/testify/suite"
)
type ClockTestSuite struct {
suite.Suite
clock Clock
wg sync.WaitGroup
ctx context.Context
ctxCancel context.CancelFunc
}
func (suite *ClockTestSuite) SetupTest() {
ctx, cancel := context.WithCancel(context.Background())
suite.ctx = ctx
suite.ctxCancel = cancel
suite.clock = Clock{
stats: &Stats{
k: StatsDefaultK,
lambda: StatsDefaultLambda,
},
tick: make(chan struct{}),
}
suite.wg.Go(func() {
suite.clock.Start(suite.ctx)
})
}
func (suite *ClockTestSuite) TearDownTest() {
suite.ctxCancel()
suite.wg.Wait()
}
func (suite *ClockTestSuite) TestTicks() {
received := 0
for range 3 {
select {
case <-suite.clock.tick:
received++
case <-time.After(2 * time.Second):
suite.Fail("timed out waiting for tick")
}
}
suite.Equal(3, received)
}
func (suite *ClockTestSuite) TestStopsOnCancel() {
select {
case <-suite.clock.tick:
case <-time.After(2 * time.Second):
suite.Fail("timed out waiting for first tick")
}
suite.ctxCancel()
time.Sleep(50 * time.Millisecond)
select {
case <-suite.clock.tick:
suite.Fail("received tick after cancel")
default:
}
}
func TestClock(t *testing.T) {
t.Parallel()
suite.Run(t, &ClockTestSuite{})
}
+44 -56
View File
@@ -4,11 +4,19 @@ import (
"bytes" "bytes"
"context" "context"
"sync" "sync"
"time"
"github.com/9seconds/mtg/v2/essentials" "github.com/9seconds/mtg/v2/essentials"
"github.com/9seconds/mtg/v2/mtglib/internal/tls" "github.com/9seconds/mtg/v2/mtglib/internal/tls"
) )
var doppelBufPool = sync.Pool{
New: func() any {
b := make([]byte, tls.MaxRecordSize)
return &b
},
}
type Conn struct { type Conn struct {
essentials.Conn essentials.Conn
@@ -18,110 +26,90 @@ type Conn struct {
type connPayload struct { type connPayload struct {
ctx context.Context ctx context.Context
ctxCancel context.CancelCauseFunc ctxCancel context.CancelCauseFunc
clock Clock stats Stats
wg sync.WaitGroup wg sync.WaitGroup
syncWriteLock sync.RWMutex
writeStream bytes.Buffer writeStream bytes.Buffer
writeCond *sync.Cond writtenCond sync.Cond
done bool
} }
func (c Conn) Write(p []byte) (int, error) { func (c Conn) Write(p []byte) (int, error) {
c.p.syncWriteLock.RLock() if len(p) == 0 {
defer c.p.syncWriteLock.RUnlock() return 0, context.Cause(c.p.ctx)
}
c.p.writeCond.L.Lock() c.p.writtenCond.L.Lock()
c.p.writeStream.Write(p) c.p.writeStream.Write(p)
c.p.writeCond.L.Unlock() c.p.writtenCond.L.Unlock()
c.p.writtenCond.Signal()
return len(p), context.Cause(c.p.ctx) return len(p), context.Cause(c.p.ctx)
} }
func (c Conn) SyncWrite(p []byte) (int, error) {
c.p.syncWriteLock.Lock()
defer c.p.syncWriteLock.Unlock()
c.p.writeCond.L.Lock()
// wait until buffer is exhausted
for c.p.writeStream.Len() != 0 && context.Cause(c.p.ctx) == nil {
c.p.writeCond.Wait()
}
c.p.writeStream.Write(p)
c.p.writeCond.L.Unlock()
if err := context.Cause(c.p.ctx); err != nil {
return len(p), err
}
c.p.writeCond.L.Lock()
// wait until data will be sent
for c.p.writeStream.Len() != 0 && context.Cause(c.p.ctx) == nil {
c.p.writeCond.Wait()
}
c.p.writeCond.L.Unlock()
return len(p), context.Cause(c.p.ctx)
}
func (c Conn) Start() {
c.p.wg.Go(func() {
c.start()
})
}
func (c Conn) start() { func (c Conn) start() {
defer c.p.writeCond.Broadcast() bp := doppelBufPool.Get().(*[]byte)
buf := *bp
defer doppelBufPool.Put(bp)
buf := [tls.MaxRecordSize]byte{} timer := time.NewTimer(c.p.stats.Delay())
defer timer.Stop()
for { for {
select { select {
case <-c.p.ctx.Done(): case <-c.p.ctx.Done():
return return
case <-c.p.clock.tick: case <-timer.C:
timer.Reset(c.p.stats.Delay())
} }
c.p.writeCond.L.Lock() size := c.p.stats.Size()
n, err := c.p.writeStream.Read(buf[:c.p.clock.stats.Size()])
c.p.writeCond.L.Unlock()
if n == 0 || err != nil { c.p.writtenCond.L.Lock()
for c.p.writeStream.Len() == 0 && !c.p.done {
c.p.writtenCond.Wait()
}
n, _ := c.p.writeStream.Read(buf[tls.SizeHeader : tls.SizeHeader+size])
c.p.writtenCond.L.Unlock()
if n == 0 {
continue continue
} }
if err := tls.WriteRecord(c.Conn, buf[:n]); err != nil { if err := tls.WriteRecordInPlace(c.Conn, buf, n); err != nil {
c.p.ctxCancel(err) c.p.ctxCancel(err)
return return
} }
c.p.writeCond.Signal()
} }
} }
func (c Conn) Stop() { func (c Conn) Stop() {
c.p.ctxCancel(nil) c.p.ctxCancel(nil)
c.p.writtenCond.L.Lock()
c.p.done = true
c.p.writtenCond.L.Unlock()
c.p.writtenCond.Broadcast()
c.p.wg.Wait() c.p.wg.Wait()
} }
func NewConn(ctx context.Context, conn essentials.Conn, stats *Stats) Conn { func NewConn(ctx context.Context, conn essentials.Conn, stats Stats) Conn {
ctx, cancel := context.WithCancelCause(ctx) ctx, cancel := context.WithCancelCause(ctx)
rv := Conn{ rv := Conn{
Conn: conn, Conn: conn,
p: &connPayload{ p: &connPayload{
ctx: ctx, ctx: ctx,
ctxCancel: cancel, ctxCancel: cancel,
writeCond: sync.NewCond(&sync.Mutex{}),
clock: Clock{
stats: stats, stats: stats,
tick: make(chan struct{}), writtenCond: sync.Cond{
L: &sync.Mutex{},
}, },
}, },
} }
rv.p.writeStream.Grow(tls.DefaultBufferSize) rv.p.writeStream.Grow(tls.DefaultBufferSize)
rv.p.wg.Go(func() {
rv.p.clock.Start(ctx)
})
rv.p.wg.Go(func() { rv.p.wg.Go(func() {
rv.start() rv.start()
}) })
+32 -131
View File
@@ -63,7 +63,7 @@ func (suite *ConnTestSuite) TearDownTest() {
} }
func (suite *ConnTestSuite) makeConn() Conn { func (suite *ConnTestSuite) makeConn() Conn {
return NewConn(suite.ctx, suite.connMock, &Stats{ return NewConn(suite.ctx, suite.connMock, Stats{
k: 2.0, k: 2.0,
lambda: 0.01, lambda: 0.01,
}) })
@@ -141,6 +141,37 @@ func (suite *ConnTestSuite) TestWriteReturnsErrorAfterStop() {
suite.Error(err) suite.Error(err)
} }
func (suite *ConnTestSuite) TestStopDoesNotDeadlockWhenStartIsWaiting() {
suite.connMock.
On("Write", mock.AnythingOfType("[]uint8")).
Return(0, nil).
Maybe()
for range 100 {
func() {
ctx, cancel := context.WithCancel(suite.ctx)
defer cancel()
c := NewConn(ctx, suite.connMock, Stats{
k: 2.0,
lambda: 0.01,
})
done := make(chan struct{})
go func() {
defer close(done)
c.Stop()
}()
select {
case <-done:
case <-time.After(2 * time.Second):
suite.Fail("Stop() deadlocked: start() likely stuck in writtenCond.Wait()")
}
}()
}
}
func (suite *ConnTestSuite) TestStopOnUnderlyingWriteError() { func (suite *ConnTestSuite) TestStopOnUnderlyingWriteError() {
suite.connMock. suite.connMock.
On("Write", mock.AnythingOfType("[]uint8")). On("Write", mock.AnythingOfType("[]uint8")).
@@ -157,136 +188,6 @@ func (suite *ConnTestSuite) TestStopOnUnderlyingWriteError() {
}, 2*time.Second, time.Millisecond) }, 2*time.Second, time.Millisecond)
} }
func (suite *ConnTestSuite) TestSyncWriteDataSent() {
suite.connMock.
On("Write", mock.AnythingOfType("[]uint8")).
Return(0, nil).
Maybe()
c := suite.makeConn()
defer c.Stop()
payload := []byte("sync hello")
n, err := c.SyncWrite(payload)
suite.NoError(err)
suite.Equal(len(payload), n)
// SyncWrite returns only after data is flushed to the wire.
assembled := &bytes.Buffer{}
reader := bytes.NewReader(suite.connMock.Written())
for {
header := make([]byte, tls.SizeHeader)
if _, err := io.ReadFull(reader, header); err != nil {
break
}
suite.Equal(byte(tls.TypeApplicationData), header[0])
length := binary.BigEndian.Uint16(header[tls.SizeRecordType+tls.SizeVersion:])
rec := make([]byte, length)
_, err := io.ReadFull(reader, rec)
suite.NoError(err)
assembled.Write(rec)
}
suite.Equal(payload, assembled.Bytes())
}
func (suite *ConnTestSuite) TestSyncWriteDrainsBufferFirst() {
suite.connMock.
On("Write", mock.AnythingOfType("[]uint8")).
Return(0, nil).
Maybe()
c := suite.makeConn()
defer c.Stop()
// Buffer some data via async Write.
_, err := c.Write([]byte("first"))
suite.NoError(err)
// SyncWrite must drain "first" before sending "second".
n, err := c.SyncWrite([]byte("second"))
suite.NoError(err)
suite.Equal(6, n)
// All data should be on the wire now.
assembled := &bytes.Buffer{}
reader := bytes.NewReader(suite.connMock.Written())
for {
header := make([]byte, tls.SizeHeader)
if _, err := io.ReadFull(reader, header); err != nil {
break
}
length := binary.BigEndian.Uint16(header[tls.SizeRecordType+tls.SizeVersion:])
rec := make([]byte, length)
_, err := io.ReadFull(reader, rec)
suite.NoError(err)
assembled.Write(rec)
}
suite.Equal([]byte("firstsecond"), assembled.Bytes())
}
func (suite *ConnTestSuite) TestSyncWriteBlocksAsyncWrite() {
suite.connMock.
On("Write", mock.AnythingOfType("[]uint8")).
Return(0, nil).
Maybe()
c := suite.makeConn()
defer c.Stop()
// Start SyncWrite — it holds exclusive lock.
syncDone := make(chan struct{})
go func() {
defer close(syncDone)
c.SyncWrite([]byte("exclusive")) //nolint: errcheck
}()
// Give SyncWrite time to acquire the lock.
time.Sleep(10 * time.Millisecond)
// Async Write should block until SyncWrite completes.
writeDone := make(chan struct{})
go func() {
defer close(writeDone)
c.Write([]byte("blocked")) //nolint: errcheck
}()
// SyncWrite should finish first.
<-syncDone
select {
case <-writeDone:
// Write completed after SyncWrite — correct.
case <-time.After(2 * time.Second):
suite.Fail("async Write did not unblock after SyncWrite completed")
}
}
func (suite *ConnTestSuite) TestSyncWriteReturnsErrorAfterStop() {
suite.connMock.
On("Write", mock.AnythingOfType("[]uint8")).
Return(0, nil).
Maybe()
c := suite.makeConn()
c.Stop()
time.Sleep(10 * time.Millisecond)
_, err := c.SyncWrite([]byte("too late"))
suite.Error(err)
}
func TestConn(t *testing.T) { func TestConn(t *testing.T) {
t.Parallel() t.Parallel()
suite.Run(t, &ConnTestSuite{}) suite.Run(t, &ConnTestSuite{})
+94 -11
View File
@@ -2,7 +2,9 @@ package doppel
import ( import (
"context" "context"
"fmt"
"sync" "sync"
"sync/atomic"
"time" "time"
"github.com/9seconds/mtg/v2/essentials" "github.com/9seconds/mtg/v2/essentials"
@@ -12,8 +14,22 @@ const (
DoppelGangerMaxDurations = 4096 DoppelGangerMaxDurations = 4096
DoppelGangerScoutRaidEach = 6 * time.Hour DoppelGangerScoutRaidEach = 6 * time.Hour
DoppelGangerScoutRepeats = 10 DoppelGangerScoutRepeats = 10
MinCertSizesToCalculate = 3
) )
// NoiseParams holds the measured cert chain size for FakeTLS noise calibration.
// If Mean is 0, the caller should use a legacy fallback.
type NoiseParams struct {
Mean int
Jitter int
}
type scoutRaidResult struct {
durations []time.Duration
certSizes []int
}
type gangerConnRequest struct { type gangerConnRequest struct {
ret chan<- Conn ret chan<- Conn
payload essentials.Conn payload essentials.Conn
@@ -31,8 +47,11 @@ type Ganger struct {
drs bool drs bool
stats *Stats stats Stats
durations []time.Duration durations []time.Duration
certSizes []int
noiseParams atomic.Pointer[NoiseParams]
connRequests chan gangerConnRequest connRequests chan gangerConnRequest
} }
@@ -48,6 +67,16 @@ func (g *Ganger) Run() {
}) })
} }
// NoiseParams returns the current cert-size-based noise parameters.
// Returns zero-value NoiseParams if not yet measured (caller should use fallback).
func (g *Ganger) NoiseParams() NoiseParams {
if p := g.noiseParams.Load(); p != nil {
return *p
}
return NoiseParams{}
}
func (g *Ganger) NewConn(conn essentials.Conn) (Conn, error) { func (g *Ganger) NewConn(conn essentials.Conn) (Conn, error) {
rvChan := make(chan Conn) rvChan := make(chan Conn)
req := gangerConnRequest{ req := gangerConnRequest{
@@ -81,10 +110,10 @@ func (g *Ganger) run() {
} }
}() }()
scoutCollectedChan := make(chan []time.Duration) scoutCollectedChan := make(chan scoutRaidResult)
currentScoutCollectedChan := scoutCollectedChan currentScoutCollectedChan := scoutCollectedChan
updatedStatsChan := make(chan *Stats) updatedStatsChan := make(chan Stats)
g.wg.Go(func() { g.wg.Go(func() {
g.runScoutRaid(scoutCollectedChan) g.runScoutRaid(scoutCollectedChan)
@@ -94,17 +123,29 @@ func (g *Ganger) run() {
select { select {
case <-g.ctx.Done(): case <-g.ctx.Done():
return return
case durations := <-currentScoutCollectedChan: case result := <-currentScoutCollectedChan:
g.durations = append(g.durations, durations...) g.durations = append(g.durations, result.durations...)
if len(g.durations) > DoppelGangerMaxDurations { if len(g.durations) > DoppelGangerMaxDurations {
g.durations = g.durations[len(g.durations)-DoppelGangerMaxDurations:] copy(g.durations, g.durations[len(g.durations)-DoppelGangerMaxDurations:])
g.durations = g.durations[:DoppelGangerMaxDurations]
}
// Update cert sizes and recompute noise params.
g.certSizes = append(g.certSizes, result.certSizes...)
if len(g.certSizes) > DoppelGangerMaxDurations {
g.certSizes = g.certSizes[len(g.certSizes)-DoppelGangerMaxDurations:]
}
if len(g.certSizes) >= MinCertSizesToCalculate {
g.updateNoiseParams()
} }
if len(g.durations) < MinDurationsToCalculate { if len(g.durations) < MinDurationsToCalculate {
continue continue
} }
durations := g.durations
currentScoutCollectedChan = nil currentScoutCollectedChan = nil
g.wg.Go(func() { g.wg.Go(func() {
select { select {
@@ -128,8 +169,45 @@ func (g *Ganger) run() {
} }
} }
func (g *Ganger) runScoutRaid(rvChan chan<- []time.Duration) { func (g *Ganger) updateNoiseParams() {
durations := []time.Duration{} if len(g.certSizes) == 0 {
return
}
sum := 0
for _, s := range g.certSizes {
sum += s
}
mean := sum / len(g.certSizes)
maxDev := 0
for _, s := range g.certSizes {
d := s - mean
if d < 0 {
d = -d
}
if d > maxDev {
maxDev = d
}
}
if maxDev < 100 {
maxDev = 100
}
np := &NoiseParams{Mean: mean, Jitter: maxDev}
g.noiseParams.Store(np)
g.logger.Info(fmt.Sprintf(
"updated noise params: mean=%d jitter=%d samples=%d",
mean, maxDev, len(g.certSizes),
))
}
func (g *Ganger) runScoutRaid(rvChan chan<- scoutRaidResult) {
var result scoutRaidResult
for range g.scoutRaidRepeats { for range g.scoutRaidRepeats {
learned, err := g.scout.Learn(g.ctx) learned, err := g.scout.Learn(g.ctx)
@@ -137,13 +215,18 @@ func (g *Ganger) runScoutRaid(rvChan chan<- []time.Duration) {
g.logger.WarningError("cannot learn", err) g.logger.WarningError("cannot learn", err)
continue continue
} }
durations = append(durations, learned...)
result.durations = append(result.durations, learned.Durations...)
if learned.CertSize > 0 {
result.certSizes = append(result.certSizes, learned.CertSize)
}
} }
select { select {
case <-g.ctx.Done(): case <-g.ctx.Done():
return return
case rvChan <- durations: case rvChan <- result:
} }
} }
@@ -173,7 +256,7 @@ func NewGanger(
scoutRaidEach: scoutEach, scoutRaidEach: scoutEach,
scoutRaidRepeats: scoutRepeats, scoutRaidRepeats: scoutRepeats,
drs: drs, drs: drs,
stats: &Stats{ stats: Stats{
k: StatsDefaultK, k: StatsDefaultK,
lambda: StatsDefaultLambda, lambda: StatsDefaultLambda,
drs: drs, drs: drs,
+56 -15
View File
@@ -12,36 +12,46 @@ import (
"github.com/9seconds/mtg/v2/mtglib/internal/tls" "github.com/9seconds/mtg/v2/mtglib/internal/tls"
) )
// ScoutResult holds measurements from a single scout HTTP request.
type ScoutResult struct {
Durations []time.Duration
CertSize int // total ApplicationData bytes during TLS handshake; 0 if unknown
}
type Scout struct { type Scout struct {
network Network network Network
urls []string urls []string
} }
func (s Scout) Learn(ctx context.Context) ([]time.Duration, error) { func (s Scout) Learn(ctx context.Context) (ScoutResult, error) {
var durations []time.Duration var combined ScoutResult
for _, url := range s.urls { for _, url := range s.urls {
learned, err := s.learn(ctx, url) learned, err := s.learn(ctx, url)
if err != nil { if err != nil {
return nil, err return ScoutResult{}, err
} }
durations = append(durations, learned...) combined.Durations = append(combined.Durations, learned.Durations...)
if learned.CertSize > 0 && combined.CertSize == 0 {
combined.CertSize = learned.CertSize
}
} }
return durations, nil return combined, nil
} }
func (s Scout) learn(ctx context.Context, url string) ([]time.Duration, error) { func (s Scout) learn(ctx context.Context, url string) (ScoutResult, error) {
client, results := s.makeClient() client, results := s.makeClient()
if !strings.HasPrefix(url, "https://") { if !strings.HasPrefix(url, "https://") {
return nil, fmt.Errorf("url %s must be https", url) return ScoutResult{}, fmt.Errorf("url %s must be https", url)
} }
req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil) req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil)
if err != nil { if err != nil {
return nil, err return ScoutResult{}, err
} }
resp, err := client.Do(req) resp, err := client.Do(req)
@@ -51,31 +61,62 @@ func (s Scout) learn(ctx context.Context, url string) ([]time.Duration, error) {
client.CloseIdleConnections() client.CloseIdleConnections()
} }
if err != nil || len(results.data) == 0 { if err != nil {
return nil, err return ScoutResult{}, err
} }
durations := []time.Duration{} data, writeIndex := results.Snapshot()
if len(data) == 0 {
return ScoutResult{}, nil
}
var result ScoutResult
// Compute inter-record durations (existing logic).
lastTimestamp := time.Time{} lastTimestamp := time.Time{}
for i, v := range results.data { for i, v := range data {
if v.recordType != tls.TypeApplicationData { if v.recordType != tls.TypeApplicationData {
continue continue
} }
if lastTimestamp.IsZero() { if lastTimestamp.IsZero() {
if i > 0 { if i > 0 {
lastTimestamp = results.data[i-1].timestamp lastTimestamp = data[i-1].timestamp
} else { } else {
lastTimestamp = v.timestamp lastTimestamp = v.timestamp
} }
} }
durations = append(durations, v.timestamp.Sub(lastTimestamp)) result.Durations = append(result.Durations, v.timestamp.Sub(lastTimestamp))
lastTimestamp = v.timestamp lastTimestamp = v.timestamp
} }
return durations, nil // Compute cert size: sum of ApplicationData payload between CCS and
// the first client Write (which marks the end of server handshake).
seenCCS := false
boundary := writeIndex
if boundary < 0 {
boundary = len(data)
}
for i, v := range data {
if i >= boundary {
break
}
if v.recordType == tls.TypeChangeCipherSpec {
seenCCS = true
continue
}
if seenCCS && v.recordType == tls.TypeApplicationData {
result.CertSize += v.payloadLen
}
}
return result, nil
} }
func (s Scout) makeClient() (*http.Client, *ScoutConnCollected) { func (s Scout) makeClient() (*http.Client, *ScoutConnCollected) {
+17 -4
View File
@@ -14,9 +14,10 @@ type ScoutConn struct {
results *ScoutConnCollected results *ScoutConnCollected
rawBuf *bytes.Buffer rawBuf *bytes.Buffer
seenCCS bool
} }
func (s ScoutConn) Read(p []byte) (int, error) { func (s *ScoutConn) Read(p []byte) (int, error) {
buf := &bytes.Buffer{} buf := &bytes.Buffer{}
for { for {
@@ -31,7 +32,11 @@ func (s ScoutConn) Read(p []byte) (int, error) {
return 0, err return 0, err
} }
s.results.Add(recordType) if recordType == tls.TypeChangeCipherSpec {
s.seenCCS = true
}
s.results.Add(recordType, int(length))
s.rawBuf.Write([]byte{recordType}) s.rawBuf.Write([]byte{recordType})
s.rawBuf.Write(tls.TLSVersion[:]) s.rawBuf.Write(tls.TLSVersion[:])
@@ -45,11 +50,19 @@ func (s ScoutConn) Read(p []byte) (int, error) {
} }
} }
func NewScoutConn(conn essentials.Conn, results *ScoutConnCollected) ScoutConn { func (s *ScoutConn) Write(p []byte) (int, error) {
if s.seenCCS {
s.results.MarkWrite()
}
return s.Conn.Write(p)
}
func NewScoutConn(conn essentials.Conn, results *ScoutConnCollected) *ScoutConn {
rawBuf := &bytes.Buffer{} rawBuf := &bytes.Buffer{}
rawBuf.Grow(tls.MaxRecordSize) rawBuf.Grow(tls.MaxRecordSize)
return ScoutConn{ return &ScoutConn{
Conn: tls.New(conn, false, false), Conn: tls.New(conn, false, false),
results: results, results: results,
rawBuf: rawBuf, rawBuf: rawBuf,
+32 -2
View File
@@ -1,6 +1,10 @@
package doppel package doppel
import "time" import (
"slices"
"sync"
"time"
)
const ( const (
ScoutConnCollectedPreallocSize = 100 ScoutConnCollectedPreallocSize = 100
@@ -9,21 +13,47 @@ const (
type ScoutConnResult struct { type ScoutConnResult struct {
timestamp time.Time timestamp time.Time
recordType byte recordType byte
payloadLen int
} }
type ScoutConnCollected struct { type ScoutConnCollected struct {
mu sync.Mutex
data []ScoutConnResult data []ScoutConnResult
writeIndex int // index at which client first wrote post-handshake data; -1 if not set
} }
func (s *ScoutConnCollected) Add(record byte) { func (s *ScoutConnCollected) Add(record byte, payloadLen int) {
s.mu.Lock()
s.data = append(s.data, ScoutConnResult{ s.data = append(s.data, ScoutConnResult{
timestamp: time.Now(), timestamp: time.Now(),
recordType: record, recordType: record,
payloadLen: payloadLen,
}) })
s.mu.Unlock()
}
// MarkWrite records the current data length as the handshake boundary.
func (s *ScoutConnCollected) MarkWrite() {
s.mu.Lock()
if s.writeIndex < 0 {
s.writeIndex = len(s.data)
}
s.mu.Unlock()
}
// Snapshot returns a copy of the collected data and the write index.
func (s *ScoutConnCollected) Snapshot() ([]ScoutConnResult, int) {
s.mu.Lock()
snapshot := slices.Clone(s.data)
writeIndex := s.writeIndex
s.mu.Unlock()
return snapshot, writeIndex
} }
func NewScoutConnCollected() *ScoutConnCollected { func NewScoutConnCollected() *ScoutConnCollected {
return &ScoutConnCollected{ return &ScoutConnCollected{
data: make([]ScoutConnResult, 0, ScoutConnCollectedPreallocSize), data: make([]ScoutConnResult, 0, ScoutConnCollectedPreallocSize),
writeIndex: -1,
} }
} }
@@ -1,6 +1,7 @@
package doppel package doppel
import ( import (
"sync"
"testing" "testing"
"time" "time"
@@ -14,28 +15,71 @@ type ScoutConnCollectedTestSuite struct {
func (suite *ScoutConnCollectedTestSuite) TestAddSingle() { func (suite *ScoutConnCollectedTestSuite) TestAddSingle() {
collected := NewScoutConnCollected() collected := NewScoutConnCollected()
collected.Add(tls.TypeApplicationData) collected.Add(tls.TypeApplicationData, 100)
suite.Len(collected.data, 1) data, _ := collected.Snapshot()
suite.Equal(byte(tls.TypeApplicationData), collected.data[0].recordType)
suite.Len(data, 1)
suite.Equal(byte(tls.TypeApplicationData), data[0].recordType)
} }
func (suite *ScoutConnCollectedTestSuite) TestAddTimestampsAreMonotonic() { func (suite *ScoutConnCollectedTestSuite) TestAddTimestampsAreMonotonic() {
collected := NewScoutConnCollected() collected := NewScoutConnCollected()
collected.Add(tls.TypeApplicationData) collected.Add(tls.TypeApplicationData, 100)
time.Sleep(time.Microsecond) time.Sleep(time.Microsecond)
collected.Add(tls.TypeApplicationData) collected.Add(tls.TypeApplicationData, 100)
time.Sleep(time.Microsecond) time.Sleep(time.Microsecond)
collected.Add(tls.TypeApplicationData) collected.Add(tls.TypeApplicationData, 100)
for i := 1; i < len(collected.data); i++ { data, _ := collected.Snapshot()
suite.True(collected.data[i].timestamp.After(collected.data[i-1].timestamp))
for i := 1; i < len(data); i++ {
suite.True(data[i].timestamp.After(data[i-1].timestamp))
} }
} }
func (suite *ScoutConnCollectedTestSuite) TestConcurrentAddSnapshot() {
collected := NewScoutConnCollected()
var wg sync.WaitGroup
wg.Add(3)
go func() {
defer wg.Done()
for i := 0; i < 1000; i++ {
collected.Add(tls.TypeApplicationData, i)
}
}()
go func() {
defer wg.Done()
for i := 0; i < 100; i++ {
collected.MarkWrite()
}
}()
go func() {
defer wg.Done()
for i := 0; i < 1000; i++ {
// call Snapshot concurrently to exercise the lock under -race
collected.Snapshot() //nolint:errcheck
}
}()
wg.Wait()
data, writeIndex := collected.Snapshot()
suite.Len(data, 1000)
suite.GreaterOrEqual(writeIndex, 0)
}
func TestScoutConnCollected(t *testing.T) { func TestScoutConnCollected(t *testing.T) {
t.Parallel() t.Parallel()
suite.Run(t, &ScoutConnCollectedTestSuite{}) suite.Run(t, &ScoutConnCollectedTestSuite{})
+2 -2
View File
@@ -22,9 +22,9 @@ func (suite *ScoutTestSuite) SetupSuite() {
} }
func (suite *ScoutTestSuite) TestCollectResults() { func (suite *ScoutTestSuite) TestCollectResults() {
durations, err := suite.scout.Learn(suite.ctx) result, err := suite.scout.Learn(suite.ctx)
suite.NoError(err) suite.NoError(err)
suite.Less(3, len(durations)) suite.Less(3, len(result.Durations))
} }
func (suite *ScoutTestSuite) TestCollectNothing() { func (suite *ScoutTestSuite) TestCollectNothing() {
+2 -2
View File
@@ -112,7 +112,7 @@ func (d *Stats) Size() int {
return TLSRecordSizeMax return TLSRecordSizeMax
} }
func NewStats(durations []time.Duration, drs bool) *Stats { func NewStats(durations []time.Duration, drs bool) Stats {
n := float64(len(durations)) n := float64(len(durations))
// in milliseconds // in milliseconds
@@ -162,7 +162,7 @@ func NewStats(durations []time.Duration, drs bool) *Stats {
// λ = (Σxᵢᵏ / n)^(1/k) // λ = (Σxᵢᵏ / n)^(1/k)
lambda := math.Pow(sumXK/n, 1.0/k) lambda := math.Pow(sumXK/n, 1.0/k)
return &Stats{ return Stats{
k: k, k: k,
lambda: lambda, lambda: lambda,
drs: drs, drs: drs,
@@ -0,0 +1,13 @@
//go:build mips || mipsle
package relay
import "github.com/9seconds/mtg/v2/mtglib/internal/tls"
const (
// MIPS is quite short in resources, and usually it means that it will run
// on Microtiks, OpenWRT-based routers or similar hardware. I think it worth
// to sacrifice a number of read syscalls (read, CPU load) to shrink
// limited RAM resources.
bufPoolSize = tls.MaxRecordPayloadSize / 2
)
@@ -0,0 +1,9 @@
//go:build !mips && !mipsle
package relay
import "github.com/9seconds/mtg/v2/mtglib/internal/tls"
const (
bufPoolSize = tls.MaxRecordPayloadSize
)
+18
View File
@@ -0,0 +1,18 @@
package relay
import "sync"
var bufPool = sync.Pool{
New: func() any {
b := make([]byte, bufPoolSize)
return &b
},
}
func acquireBuffer() *[]byte {
return bufPool.Get().(*[]byte)
}
func releaseBuffer(p *[]byte) {
bufPool.Put(p)
}
+6 -6
View File
@@ -6,7 +6,6 @@ import (
"io" "io"
"github.com/9seconds/mtg/v2/essentials" "github.com/9seconds/mtg/v2/essentials"
"github.com/9seconds/mtg/v2/mtglib/internal/tls"
) )
func Relay(ctx context.Context, log Logger, telegramConn, clientConn essentials.Conn) { func Relay(ctx context.Context, log Logger, telegramConn, clientConn essentials.Conn) {
@@ -16,11 +15,11 @@ func Relay(ctx context.Context, log Logger, telegramConn, clientConn essentials.
ctx, cancel := context.WithCancel(ctx) ctx, cancel := context.WithCancel(ctx)
defer cancel() defer cancel()
go func() { stop := context.AfterFunc(ctx, func() {
<-ctx.Done()
telegramConn.Close() //nolint: errcheck telegramConn.Close() //nolint: errcheck
clientConn.Close() //nolint: errcheck clientConn.Close() //nolint: errcheck
}() })
defer stop()
closeChan := make(chan struct{}) closeChan := make(chan struct{})
@@ -36,12 +35,13 @@ func Relay(ctx context.Context, log Logger, telegramConn, clientConn essentials.
} }
func pump(log Logger, src, dst essentials.Conn, direction string) { func pump(log Logger, src, dst essentials.Conn, direction string) {
var buf [tls.MaxRecordPayloadSize]byte buf := acquireBuffer()
defer releaseBuffer(buf)
defer src.CloseRead() //nolint: errcheck defer src.CloseRead() //nolint: errcheck
defer dst.CloseWrite() //nolint: errcheck defer dst.CloseWrite() //nolint: errcheck
n, err := io.CopyBuffer(src, dst, buf[:]) n, err := io.CopyBuffer(src, dst, *buf)
switch { switch {
case err == nil: case err == nil:
-2
View File
@@ -34,7 +34,6 @@ type Conn struct {
type connPayload struct { type connPayload struct {
readBuf bytes.Buffer readBuf bytes.Buffer
writeBuf bytes.Buffer
connBuffered *bufio.Reader connBuffered *bufio.Reader
read bool read bool
write bool write bool
@@ -80,7 +79,6 @@ func New(conn essentials.Conn, read, write bool) Conn {
} }
newConn.p.readBuf.Grow(DefaultBufferSize) newConn.p.readBuf.Grow(DefaultBufferSize)
newConn.p.writeBuf.Grow(DefaultBufferSize)
return newConn return newConn
} }
+32 -80
View File
@@ -6,13 +6,12 @@ import (
"crypto/sha256" "crypto/sha256"
"crypto/subtle" "crypto/subtle"
"encoding/binary" "encoding/binary"
"errors"
"fmt" "fmt"
"io" "io"
"net" "net"
"slices" "slices"
"time" "time"
"github.com/9seconds/mtg/v2/mtglib/internal/tls"
) )
const ( const (
@@ -22,12 +21,19 @@ const (
// record_type(1) + version(2) + size(2) + handshake_type(1) + uint24_length(3) + client_version(2) // record_type(1) + version(2) + size(2) + handshake_type(1) + uint24_length(3) + client_version(2)
RandomOffset = 1 + 2 + 2 + 1 + 3 + 2 RandomOffset = 1 + 2 + 2 + 1 + 3 + 2
// https://datatracker.ietf.org/doc/html/rfc8701#name-grease-values
// https://medium.com/asecuritysite-when-bob-met-alice/in-cybersecurity-what-is-grease-9f8850558dea
GreaseMask = 0x0f0f
GreaseValueType = 0x0a0a
sniDNSNamesListType = 0 sniDNSNamesListType = 0
) )
var ( var (
emptyRandom = [RandomLen]byte{} emptyRandom = [RandomLen]byte{}
extTypeSNI = [2]byte{} extTypeSNI = [2]byte{}
ErrCannotFindCipher = errors.New("cannot find a cipher")
) )
type ClientHello struct { type ClientHello struct {
@@ -42,11 +48,6 @@ func ReadClientHello(
hostname string, hostname string,
tolerateTimeSkewness time.Duration, tolerateTimeSkewness time.Duration,
) (*ClientHello, error) { ) (*ClientHello, error) {
if err := conn.SetReadDeadline(time.Now().Add(ClientHelloReadTimeout)); err != nil {
return nil, fmt.Errorf("cannot set read deadline: %w", err)
}
defer conn.SetReadDeadline(resetDeadline) //nolint: errcheck
// This is how FakeTLS is organized: // This is how FakeTLS is organized:
// 1. We create sha256 HMAC with a given secret // 1. We create sha256 HMAC with a given secret
// 2. We dump there a whole TLS frame except of the fact that random // 2. We dump there a whole TLS frame except of the fact that random
@@ -56,25 +57,17 @@ func ReadClientHello(
// 4. New digest should be all 0 except of last 4 bytes // 4. New digest should be all 0 except of last 4 bytes
// 5. Last 4 bytes are little endian uint32 of UNIX timestamp when // 5. Last 4 bytes are little endian uint32 of UNIX timestamp when
// this message was created. // this message was created.
handshakeCopyBuf := &bytes.Buffer{} clientHelloCopy, handshakeReader, err := parseClientHello(conn)
reader := io.TeeReader(conn, handshakeCopyBuf)
reader, err := parseTLSHeader(reader)
if err != nil { if err != nil {
return nil, fmt.Errorf("cannot parse tls header: %w", err) return nil, fmt.Errorf("cannot read client hello: %w", err)
} }
reader, err = parseHandshakeHeader(reader) hello, err := parseHandshake(handshakeReader)
if err != nil {
return nil, fmt.Errorf("cannot parse handshake header: %w", err)
}
hello, err := parseHandshake(reader)
if err != nil { if err != nil {
return nil, fmt.Errorf("cannot parse handshake: %w", err) return nil, fmt.Errorf("cannot parse handshake: %w", err)
} }
sniHostnames, err := parseSNI(reader) sniHostnames, err := parseSNI(handshakeReader)
if err != nil { if err != nil {
return nil, fmt.Errorf("cannot parse SNI: %w", err) return nil, fmt.Errorf("cannot parse SNI: %w", err)
} }
@@ -85,10 +78,10 @@ func ReadClientHello(
digest := hmac.New(sha256.New, secret) digest := hmac.New(sha256.New, secret)
// we write a copy of the handshake with client random all nullified. // we write a copy of the handshake with client random all nullified.
digest.Write(handshakeCopyBuf.Next(RandomOffset)) digest.Write(clientHelloCopy.Next(RandomOffset))
handshakeCopyBuf.Next(RandomLen) clientHelloCopy.Next(RandomLen)
digest.Write(emptyRandom[:]) digest.Write(emptyRandom[:])
digest.Write(handshakeCopyBuf.Bytes()) digest.Write(clientHelloCopy.Bytes())
computed := digest.Sum(nil) computed := digest.Sum(nil)
@@ -110,58 +103,6 @@ func ReadClientHello(
return hello, nil return hello, nil
} }
func parseTLSHeader(r io.Reader) (io.Reader, error) {
// record_type(1) + version(2) + size(2)
// 16 - type is 0x16 (handshake record)
// 03 01 - protocol version is "3,1" (also known as TLS 1.0)
// 00 f8 - 0xF8 (248) bytes of handshake message follows
header := [1 + 2 + 2]byte{}
if _, err := io.ReadFull(r, header[:]); err != nil {
return nil, fmt.Errorf("cannot read record header: %w", err)
}
if header[0] != tls.TypeHandshake {
return nil, fmt.Errorf("unexpected record type %#x", header[0])
}
if header[1] != 3 || header[2] != 1 {
return nil, fmt.Errorf("unexpected protocol version %#x %#x", header[1], header[2])
}
length := int64(binary.BigEndian.Uint16(header[3:]))
buf := &bytes.Buffer{}
_, err := io.CopyN(buf, r, length)
return buf, err
}
func parseHandshakeHeader(r io.Reader) (io.Reader, error) {
// type(1) + size(3 / uint24)
// 01 - handshake message type 0x01 (client hello)
// 00 00 f4 - 0xF4 (244) bytes of client hello data follows
header := [1 + 3]byte{}
if _, err := io.ReadFull(r, header[:]); err != nil {
return nil, fmt.Errorf("cannot read handshake header: %w", err)
}
if header[0] != TypeHandshakeClient {
return nil, fmt.Errorf("incorrect handshake type: %#x", header[0])
}
// unfortunately there is not uint24 in golang, so we just reust header
header[0] = 0
length := int64(binary.BigEndian.Uint32(header[:]))
buf := &bytes.Buffer{}
_, err := io.CopyN(buf, r, length)
return buf, err
}
func parseHandshake(r io.Reader) (*ClientHello, error) { func parseHandshake(r io.Reader) (*ClientHello, error) {
// A protocol version of "3,3" (meaning TLS 1.2) is given. // A protocol version of "3,3" (meaning TLS 1.2) is given.
header := [2]byte{} header := [2]byte{}
@@ -192,16 +133,27 @@ func parseHandshake(r io.Reader) (*ClientHello, error) {
cipherSuiteLen := int64(binary.BigEndian.Uint16(header[:])) cipherSuiteLen := int64(binary.BigEndian.Uint16(header[:]))
// we do not care about picking up any cipher. we pick the first one, // Pick the first non-GREASE cipher suite from the list.
// so it is always should be present. // Real TLS servers never select GREASE values (RFC 8701, pattern 0x?a?a),
// so echoing them back is a trivial DPI fingerprint.
// cipherSuiteLen is in bytes; each cipher suite is 2 bytes.
for range cipherSuiteLen / 2 {
if _, err := io.ReadFull(r, header[:]); err != nil { if _, err := io.ReadFull(r, header[:]); err != nil {
return nil, fmt.Errorf("cannot read first cipher suite: %w", err) return nil, fmt.Errorf("cannot read cipher suite: %w", err)
} }
hello.CipherSuite = binary.BigEndian.Uint16(header[:]) if hello.CipherSuite != 0 {
// do not forget we have to scan until the end
continue
}
if _, err := io.CopyN(io.Discard, r, cipherSuiteLen-2); err != nil { if cs := binary.BigEndian.Uint16(header[:]); cs&GreaseMask != GreaseValueType {
return nil, fmt.Errorf("cannot skip remaining cipher suites: %w", err) hello.CipherSuite = cs
}
}
if hello.CipherSuite == 0 {
return nil, ErrCannotFindCipher
} }
if _, err := io.ReadFull(r, header[:1]); err != nil { if _, err := io.ReadFull(r, header[:1]); err != nil {
@@ -12,7 +12,6 @@ import (
"github.com/9seconds/mtg/v2/mtglib" "github.com/9seconds/mtg/v2/mtglib"
"github.com/9seconds/mtg/v2/mtglib/internal/tls/fake" "github.com/9seconds/mtg/v2/mtglib/internal/tls/fake"
"github.com/stretchr/testify/assert" "github.com/stretchr/testify/assert"
"github.com/stretchr/testify/mock"
"github.com/stretchr/testify/require" "github.com/stretchr/testify/require"
"github.com/stretchr/testify/suite" "github.com/stretchr/testify/suite"
) )
@@ -71,11 +70,6 @@ func (suite *ParseClientHelloSnapshotTestSuite) makeConn(data []byte) *parseClie
readBuf: readBuf, readBuf: readBuf,
} }
connMock.
On("SetReadDeadline", mock.AnythingOfType("time.Time")).
Twice().
Return(nil)
return connMock return connMock
} }
+290 -23
View File
@@ -2,9 +2,11 @@ package fake_test
import ( import (
"bytes" "bytes"
cryptotls "crypto/tls"
"encoding/binary" "encoding/binary"
"errors" "encoding/json"
"io" "io"
"os"
"testing" "testing"
"time" "time"
@@ -12,7 +14,6 @@ import (
"github.com/9seconds/mtg/v2/mtglib" "github.com/9seconds/mtg/v2/mtglib"
"github.com/9seconds/mtg/v2/mtglib/internal/tls" "github.com/9seconds/mtg/v2/mtglib/internal/tls"
"github.com/9seconds/mtg/v2/mtglib/internal/tls/fake" "github.com/9seconds/mtg/v2/mtglib/internal/tls/fake"
"github.com/stretchr/testify/mock"
"github.com/stretchr/testify/require" "github.com/stretchr/testify/require"
"github.com/stretchr/testify/suite" "github.com/stretchr/testify/suite"
) )
@@ -51,11 +52,6 @@ func (suite *ParseClientHelloTestSuite) SetupTest() {
suite.connMock = &parseClientHelloConnMock{ suite.connMock = &parseClientHelloConnMock{
readBuf: suite.readBuf, readBuf: suite.readBuf,
} }
suite.connMock.
On("SetReadDeadline", mock.AnythingOfType("time.Time")).
Twice().
Return(nil)
} }
func (suite *ParseClientHelloTestSuite) TearDownTest() { func (suite *ParseClientHelloTestSuite) TearDownTest() {
@@ -67,23 +63,11 @@ type ParseClientHello_TLSHeaderTestSuite struct {
} }
func (suite *ParseClientHello_TLSHeaderTestSuite) TestEmpty() { func (suite *ParseClientHello_TLSHeaderTestSuite) TestEmpty() {
suite.connMock.ExpectedCalls = []*mock.Call{}
suite.connMock.
On("SetReadDeadline", mock.AnythingOfType("time.Time")).
Once().
Return(errors.New("fail"))
_, err := fake.ReadClientHello(suite.connMock, suite.secret.Key[:], suite.secret.Host, TolerateTime) _, err := fake.ReadClientHello(suite.connMock, suite.secret.Key[:], suite.secret.Host, TolerateTime)
suite.ErrorContains(err, "fail") suite.ErrorContains(err, "cannot read client hello")
} }
func (suite *ParseClientHello_TLSHeaderTestSuite) TestNothing() { func (suite *ParseClientHello_TLSHeaderTestSuite) TestNothing() {
suite.connMock.ExpectedCalls = []*mock.Call{}
suite.connMock.
On("SetReadDeadline", mock.AnythingOfType("time.Time")).
Twice().
Return(nil)
_, err := fake.ReadClientHello(suite.connMock, suite.secret.Key[:], suite.secret.Host, TolerateTime) _, err := fake.ReadClientHello(suite.connMock, suite.secret.Key[:], suite.secret.Host, TolerateTime)
suite.ErrorIs(err, io.EOF) suite.ErrorIs(err, io.EOF)
} }
@@ -232,12 +216,13 @@ func (suite *ParseClientHelloHandshakeBodyTestSuite) TestCannotReadCipherSuiteLe
} }
func (suite *ParseClientHelloHandshakeBodyTestSuite) TestCannotReadFirstCipherSuite() { func (suite *ParseClientHelloHandshakeBodyTestSuite) TestCannotReadFirstCipherSuite() {
body := make([]byte, 2+fake.RandomLen+1+2) body := make([]byte, 2+fake.RandomLen+1+2+1) // cipherSuiteLen=2 but only 1 byte available
binary.BigEndian.PutUint16(body[2+fake.RandomLen+1:], 2)
suite.writeBody(body) suite.writeBody(body)
_, err := fake.ReadClientHello(suite.connMock, suite.secret.Key[:], suite.secret.Host, TolerateTime) _, err := fake.ReadClientHello(suite.connMock, suite.secret.Key[:], suite.secret.Host, TolerateTime)
suite.ErrorContains(err, "cannot read first cipher suite") suite.ErrorContains(err, "cannot read cipher suite")
} }
func (suite *ParseClientHelloHandshakeBodyTestSuite) TestCannotSkipRemainingCipherSuites() { func (suite *ParseClientHelloHandshakeBodyTestSuite) TestCannotSkipRemainingCipherSuites() {
@@ -247,12 +232,27 @@ func (suite *ParseClientHelloHandshakeBodyTestSuite) TestCannotSkipRemainingCiph
suite.writeBody(body) suite.writeBody(body)
_, err := fake.ReadClientHello(suite.connMock, suite.secret.Key[:], suite.secret.Host, TolerateTime) _, err := fake.ReadClientHello(suite.connMock, suite.secret.Key[:], suite.secret.Host, TolerateTime)
suite.ErrorContains(err, "cannot skip remaining cipher suites") suite.ErrorContains(err, "cannot read cipher suite")
}
func (suite *ParseClientHelloHandshakeBodyTestSuite) TestCannotFindCipher() {
// All cipher suites are GREASE values — must return ErrCannotFindCipher.
body := make([]byte, 2+fake.RandomLen+1+2+4+1)
binary.BigEndian.PutUint16(body[2+fake.RandomLen+1:], 4)
binary.BigEndian.PutUint16(body[2+fake.RandomLen+1+2:], 0x0a0a)
binary.BigEndian.PutUint16(body[2+fake.RandomLen+1+2+2:], 0x1a1a)
body[2+fake.RandomLen+1+2+4] = 1
suite.writeBody(body)
_, err := fake.ReadClientHello(suite.connMock, suite.secret.Key[:], suite.secret.Host, TolerateTime)
suite.ErrorIs(err, fake.ErrCannotFindCipher)
} }
func (suite *ParseClientHelloHandshakeBodyTestSuite) TestCannotReadCompressionMethodsLength() { func (suite *ParseClientHelloHandshakeBodyTestSuite) TestCannotReadCompressionMethodsLength() {
body := make([]byte, 2+fake.RandomLen+1+2+2) body := make([]byte, 2+fake.RandomLen+1+2+2)
binary.BigEndian.PutUint16(body[2+fake.RandomLen+1:], 2) binary.BigEndian.PutUint16(body[2+fake.RandomLen+1:], 2)
binary.BigEndian.PutUint16(body[2+fake.RandomLen+1+2:], cryptotls.TLS_AES_128_GCM_SHA256)
suite.writeBody(body) suite.writeBody(body)
@@ -263,6 +263,7 @@ func (suite *ParseClientHelloHandshakeBodyTestSuite) TestCannotReadCompressionMe
func (suite *ParseClientHelloHandshakeBodyTestSuite) TestCannotSkipCompressionMethods() { func (suite *ParseClientHelloHandshakeBodyTestSuite) TestCannotSkipCompressionMethods() {
body := make([]byte, 2+fake.RandomLen+1+2+2+1) body := make([]byte, 2+fake.RandomLen+1+2+2+1)
binary.BigEndian.PutUint16(body[2+fake.RandomLen+1:], 2) binary.BigEndian.PutUint16(body[2+fake.RandomLen+1:], 2)
binary.BigEndian.PutUint16(body[2+fake.RandomLen+1+2:], cryptotls.TLS_AES_128_GCM_SHA256)
body[2+fake.RandomLen+1+2+2] = 1 body[2+fake.RandomLen+1+2+2] = 1
suite.writeBody(body) suite.writeBody(body)
@@ -298,6 +299,7 @@ func (suite *ParseClientHelloSNITestSuite) writeExtensions(extensions []byte) {
// cipherSuite(2) + compressionLen(1) + compression(1) = 41 // cipherSuite(2) + compressionLen(1) + compression(1) = 41
body := make([]byte, 41) body := make([]byte, 41)
binary.BigEndian.PutUint16(body[35:], 2) binary.BigEndian.PutUint16(body[35:], 2)
binary.BigEndian.PutUint16(body[37:], cryptotls.TLS_AES_128_GCM_SHA256)
body[39] = 1 body[39] = 1
suite.readBuf.Write(body) suite.readBuf.Write(body)
@@ -393,3 +395,268 @@ func TestParseClientHelloSNI(t *testing.T) {
t.Parallel() t.Parallel()
suite.Run(t, &ParseClientHelloSNITestSuite{}) suite.Run(t, &ParseClientHelloSNITestSuite{})
} }
// fragmentTLSRecord splits a single TLS record into n TLS records by
// dividing the payload into roughly equal parts. Each part gets its own
// TLS record header with the same record type and version.
func fragmentTLSRecord(t testing.TB, full []byte, n int) []byte {
t.Helper()
recordType := full[0]
version := full[1:3]
payload := full[tls.SizeHeader:]
chunkSize := len(payload) / n
result := &bytes.Buffer{}
for i := 0; i < n; i++ {
start := i * chunkSize
end := start + chunkSize
if i == n-1 {
end = len(payload)
}
chunk := payload[start:end]
result.WriteByte(recordType)
result.Write(version)
require.NoError(t, binary.Write(result, binary.BigEndian, uint16(len(chunk))))
result.Write(chunk)
}
return result.Bytes()
}
// splitPayloadAt creates two TLS records from a single record by splitting
// the payload at the given byte position.
func splitPayloadAt(t testing.TB, full []byte, pos int) []byte {
t.Helper()
payload := full[tls.SizeHeader:]
buf := &bytes.Buffer{}
buf.WriteByte(tls.TypeHandshake)
buf.Write(full[1:3])
require.NoError(t, binary.Write(buf, binary.BigEndian, uint16(pos)))
buf.Write(payload[:pos])
buf.WriteByte(tls.TypeHandshake)
buf.Write(full[1:3])
require.NoError(t, binary.Write(buf, binary.BigEndian, uint16(len(payload)-pos)))
buf.Write(payload[pos:])
return buf.Bytes()
}
type ParseClientHelloFragmentedTestSuite struct {
suite.Suite
secret mtglib.Secret
snapshot *clientHelloSnapshot
}
func (s *ParseClientHelloFragmentedTestSuite) SetupSuite() {
parsed, err := mtglib.ParseSecret(
"ee367a189aee18fa31c190054efd4a8e9573746f726167652e676f6f676c65617069732e636f6d",
)
require.NoError(s.T(), err)
s.secret = parsed
fileData, err := os.ReadFile("testdata/client-hello-ok-19dfe38384b9884b.json")
require.NoError(s.T(), err)
s.snapshot = &clientHelloSnapshot{}
require.NoError(s.T(), json.Unmarshal(fileData, s.snapshot))
}
func (s *ParseClientHelloFragmentedTestSuite) makeConn(data []byte) *parseClientHelloConnMock {
readBuf := &bytes.Buffer{}
readBuf.Write(data)
connMock := &parseClientHelloConnMock{
readBuf: readBuf,
}
return connMock
}
func (s *ParseClientHelloFragmentedTestSuite) TestReassemblySuccess() {
full := s.snapshot.GetFull()
tests := []struct {
name string
data []byte
}{
{"two equal fragments", fragmentTLSRecord(s.T(), full, 2)},
{"three equal fragments", fragmentTLSRecord(s.T(), full, 3)},
{"single byte first fragment", splitPayloadAt(s.T(), full, 1)},
{"three byte first fragment", splitPayloadAt(s.T(), full, 3)},
}
for _, tt := range tests {
s.Run(tt.name, func() {
connMock := s.makeConn(tt.data)
defer connMock.AssertExpectations(s.T())
hello, err := fake.ReadClientHello(
connMock,
s.secret.Key[:],
s.secret.Host,
TolerateTime,
)
s.Require().NoError(err)
s.Equal(s.snapshot.GetRandom(), hello.Random[:])
s.Equal(s.snapshot.GetSessionID(), hello.SessionID)
s.Equal(uint16(s.snapshot.CipherSuite), hello.CipherSuite)
})
}
}
func (s *ParseClientHelloFragmentedTestSuite) TestReassemblyErrors() {
full := s.snapshot.GetFull()
payload := full[tls.SizeHeader:]
tests := []struct {
name string
buildData func() []byte
errMsg string
}{
{
name: "wrong continuation record type",
buildData: func() []byte {
buf := &bytes.Buffer{}
buf.WriteByte(tls.TypeHandshake)
buf.Write(full[1:3])
require.NoError(s.T(), binary.Write(buf, binary.BigEndian, uint16(10)))
buf.Write(payload[:10])
// Wrong type: application data instead of handshake
buf.WriteByte(tls.TypeApplicationData)
buf.Write(full[1:3])
require.NoError(s.T(), binary.Write(buf, binary.BigEndian, uint16(len(payload)-10)))
buf.Write(payload[10:])
return buf.Bytes()
},
errMsg: "unexpected record type",
},
{
name: "too many continuation records",
buildData: func() []byte {
// Handshake header claiming 256 bytes, but we only send 1 byte per continuation
handshakePayload := []byte{0x01, 0x00, 0x01, 0x00}
buf := &bytes.Buffer{}
buf.WriteByte(tls.TypeHandshake)
buf.Write([]byte{3, 1})
require.NoError(s.T(), binary.Write(buf, binary.BigEndian, uint16(len(handshakePayload))))
buf.Write(handshakePayload)
for range 11 {
buf.WriteByte(tls.TypeHandshake)
buf.Write([]byte{3, 1})
require.NoError(s.T(), binary.Write(buf, binary.BigEndian, uint16(1)))
buf.WriteByte(0xAB)
}
return buf.Bytes()
},
errMsg: "too many fragments",
},
{
name: "zero-length continuation record",
buildData: func() []byte {
buf := &bytes.Buffer{}
buf.WriteByte(tls.TypeHandshake)
buf.Write(full[1:3])
require.NoError(s.T(), binary.Write(buf, binary.BigEndian, uint16(10)))
buf.Write(payload[:10])
// Valid header but zero-length payload
buf.WriteByte(tls.TypeHandshake)
buf.Write(full[1:3])
require.NoError(s.T(), binary.Write(buf, binary.BigEndian, uint16(0)))
return buf.Bytes()
},
errMsg: "cannot read record header",
},
{
name: "wrong continuation record version",
buildData: func() []byte {
buf := &bytes.Buffer{}
buf.WriteByte(tls.TypeHandshake)
buf.Write(full[1:3])
require.NoError(s.T(), binary.Write(buf, binary.BigEndian, uint16(10)))
buf.Write(payload[:10])
// Wrong version: 3.3 instead of 3.1
buf.WriteByte(tls.TypeHandshake)
buf.Write([]byte{3, 3})
require.NoError(s.T(), binary.Write(buf, binary.BigEndian, uint16(len(payload)-10)))
buf.Write(payload[10:])
return buf.Bytes()
},
errMsg: "unexpected protocol version",
},
{
name: "handshake message too large",
buildData: func() []byte {
// Handshake header claiming 0x010000 (65536) bytes — exceeds 0xFFFF limit
handshakePayload := []byte{0x01, 0x01, 0x00, 0x00}
buf := &bytes.Buffer{}
buf.WriteByte(tls.TypeHandshake)
buf.Write([]byte{3, 1})
require.NoError(s.T(), binary.Write(buf, binary.BigEndian, uint16(len(handshakePayload))))
buf.Write(handshakePayload)
return buf.Bytes()
},
errMsg: "cannot read record header",
},
{
name: "truncated continuation record header",
buildData: func() []byte {
buf := &bytes.Buffer{}
buf.WriteByte(tls.TypeHandshake)
buf.Write(full[1:3])
require.NoError(s.T(), binary.Write(buf, binary.BigEndian, uint16(10)))
buf.Write(payload[:10])
// Connection ends mid-header (only 2 bytes)
buf.WriteByte(tls.TypeHandshake)
buf.WriteByte(3)
return buf.Bytes()
},
errMsg: "cannot read record header",
},
{
name: "truncated continuation record payload",
buildData: func() []byte {
buf := &bytes.Buffer{}
buf.WriteByte(tls.TypeHandshake)
buf.Write(full[1:3])
require.NoError(s.T(), binary.Write(buf, binary.BigEndian, uint16(10)))
buf.Write(payload[:10])
// Claims 100 bytes but no payload follows
buf.WriteByte(tls.TypeHandshake)
buf.Write(full[1:3])
require.NoError(s.T(), binary.Write(buf, binary.BigEndian, uint16(100)))
return buf.Bytes()
},
errMsg: "EOF",
},
}
for _, tt := range tests {
s.Run(tt.name, func() {
connMock := s.makeConn(tt.buildData())
defer connMock.AssertExpectations(s.T())
_, err := fake.ReadClientHello(
connMock,
s.secret.Key[:],
s.secret.Host,
TolerateTime,
)
s.ErrorContains(err, tt.errMsg)
})
}
}
func TestParseClientHelloFragmented(t *testing.T) {
t.Parallel()
suite.Run(t, &ParseClientHelloFragmentedTestSuite{})
}
+1 -10
View File
@@ -2,15 +2,6 @@ package fake
import ( import (
"errors" "errors"
"time"
) )
const ( var ErrBadDigest = errors.New("incorrect client random")
ClientHelloReadTimeout = 5 * time.Second
)
var (
resetDeadline time.Time
ErrBadDigest = errors.New("incorrect client random")
)
+33 -6
View File
@@ -13,6 +13,14 @@ import (
"golang.org/x/crypto/curve25519" "golang.org/x/crypto/curve25519"
) )
// NoiseParams controls the size of the fake ApplicationData record
// in ServerHello. If Mean is 0, the legacy random range (2500-4700)
// is used.
type NoiseParams struct {
Mean int
Jitter int
}
const ( const (
TypeHandshakeServer = 0x02 TypeHandshakeServer = 0x02
ChangeCipherValue = 0x01 ChangeCipherValue = 0x01
@@ -32,13 +40,13 @@ var serverHelloSuffix = []byte{
0x00, 0x20, // 32 bytes of key 0x00, 0x20, // 32 bytes of key
} }
func SendServerHello(w io.Writer, secret []byte, clientHello *ClientHello) error { func SendServerHello(w io.Writer, secret []byte, clientHello *ClientHello, noise NoiseParams) error {
buf := &bytes.Buffer{} buf := &bytes.Buffer{}
buf.Grow(tls.MaxRecordSize) buf.Grow(tls.MaxRecordSize)
generateServerHello(buf, clientHello) generateServerHello(buf, clientHello)
generateChangeCipherValue(buf) generateChangeCipherValue(buf)
generateNoise(buf) generateNoise(buf, noise)
packet := buf.Bytes() packet := buf.Bytes()
digest := hmac.New(sha256.New, secret) digest := hmac.New(sha256.New, secret)
@@ -124,12 +132,31 @@ func generateChangeCipherValue(buf *bytes.Buffer) {
buf.WriteByte(ChangeCipherValue) buf.WriteByte(ChangeCipherValue)
} }
func generateNoise(buf *bytes.Buffer) { // generateNoise writes a single ApplicationData record mimicking the combined
data := make([]byte, int64(1024+rnd.IntN(3092))) // size of a real TLS 1.3 encrypted server handshake (EncryptedExtensions +
// Certificate chain + CertificateVerify + Finished).
//
// NOTE: Must be exactly ONE ApplicationData record — the Telegram client reads
// ServerHello + CCS + 1 ApplicationData and computes HMAC over all three.
// Multiple records would cause HMAC mismatch and connection failure.
func generateNoise(buf *bytes.Buffer, noise NoiseParams) {
var size int
if _, err := rand.Read(data[:]); err != nil { if noise.Mean > 0 && noise.Jitter > 0 {
// Calibrated: use measured cert chain size ± jitter.
size = noise.Mean - noise.Jitter + rnd.IntN(2*noise.Jitter)
if size < 1000 {
size = 1000
}
} else {
// Legacy fallback: random in 2500-4700 range.
size = 2500 + rnd.IntN(2200)
}
data := make([]byte, size)
if _, err := rand.Read(data); err != nil {
panic(err) panic(err)
} }
tls.WriteRecord(buf, data[:]) //nolint: errcheck tls.WriteRecord(buf, data) //nolint: errcheck
} }
+32 -6
View File
@@ -8,7 +8,6 @@ import (
"testing" "testing"
"github.com/9seconds/mtg/v2/mtglib" "github.com/9seconds/mtg/v2/mtglib"
"github.com/9seconds/mtg/v2/mtglib/internal/doppel"
"github.com/9seconds/mtg/v2/mtglib/internal/tls" "github.com/9seconds/mtg/v2/mtglib/internal/tls"
"github.com/9seconds/mtg/v2/mtglib/internal/tls/fake" "github.com/9seconds/mtg/v2/mtglib/internal/tls/fake"
"github.com/stretchr/testify/suite" "github.com/stretchr/testify/suite"
@@ -39,7 +38,7 @@ func (suite *SendServerHelloTestSuite) SetupTest() {
} }
func (suite *SendServerHelloTestSuite) TestRecordStructure() { func (suite *SendServerHelloTestSuite) TestRecordStructure() {
err := fake.SendServerHello(suite.buf, suite.secret.Key[:], suite.hello) err := fake.SendServerHello(suite.buf, suite.secret.Key[:], suite.hello, fake.NoiseParams{})
suite.NoError(err) suite.NoError(err)
var rec bytes.Buffer var rec bytes.Buffer
@@ -59,13 +58,13 @@ func (suite *SendServerHelloTestSuite) TestRecordStructure() {
recordType, length, err := tls.ReadRecord(suite.buf, &rec) recordType, length, err := tls.ReadRecord(suite.buf, &rec)
suite.NoError(err) suite.NoError(err)
suite.Equal(byte(tls.TypeApplicationData), recordType) suite.Equal(byte(tls.TypeApplicationData), recordType)
suite.Greater(length, int64(doppel.TLSRecordSizeStart)) suite.GreaterOrEqual(length, int64(2500))
suite.Empty(suite.buf.Bytes()) suite.Empty(suite.buf.Bytes())
} }
func (suite *SendServerHelloTestSuite) TestHMAC() { func (suite *SendServerHelloTestSuite) TestHMAC() {
err := fake.SendServerHello(suite.buf, suite.secret.Key[:], suite.hello) err := fake.SendServerHello(suite.buf, suite.secret.Key[:], suite.hello, fake.NoiseParams{})
suite.NoError(err) suite.NoError(err)
packet := make([]byte, suite.buf.Len()) packet := make([]byte, suite.buf.Len())
@@ -83,7 +82,7 @@ func (suite *SendServerHelloTestSuite) TestHMAC() {
} }
func (suite *SendServerHelloTestSuite) TestHandshakePayload() { func (suite *SendServerHelloTestSuite) TestHandshakePayload() {
err := fake.SendServerHello(suite.buf, suite.secret.Key[:], suite.hello) err := fake.SendServerHello(suite.buf, suite.secret.Key[:], suite.hello, fake.NoiseParams{})
suite.NoError(err) suite.NoError(err)
packet := suite.buf.Bytes() packet := suite.buf.Bytes()
@@ -105,7 +104,7 @@ func (suite *SendServerHelloTestSuite) TestHandshakePayload() {
} }
func (suite *SendServerHelloTestSuite) TestChangeCipherSpec() { func (suite *SendServerHelloTestSuite) TestChangeCipherSpec() {
err := fake.SendServerHello(suite.buf, suite.secret.Key[:], suite.hello) err := fake.SendServerHello(suite.buf, suite.secret.Key[:], suite.hello, fake.NoiseParams{})
suite.NoError(err) suite.NoError(err)
// Skip first record // Skip first record
@@ -124,6 +123,33 @@ func (suite *SendServerHelloTestSuite) TestChangeCipherSpec() {
suite.Equal([]byte{fake.ChangeCipherValue}, rec.Bytes()) suite.Equal([]byte{fake.ChangeCipherValue}, rec.Bytes())
} }
func (suite *SendServerHelloTestSuite) TestCalibratedNoiseSize() {
noise := fake.NoiseParams{Mean: 6480, Jitter: 100}
err := fake.SendServerHello(suite.buf, suite.secret.Key[:], suite.hello, noise)
suite.NoError(err)
var rec bytes.Buffer
// Skip ServerHello
_, _, err = tls.ReadRecord(suite.buf, &rec)
suite.NoError(err)
// Skip ChangeCipherSpec
rec.Reset()
_, _, err = tls.ReadRecord(suite.buf, &rec)
suite.NoError(err)
// Read noise ApplicationData
rec.Reset()
recordType, length, err := tls.ReadRecord(suite.buf, &rec)
suite.NoError(err)
suite.Equal(byte(tls.TypeApplicationData), recordType)
// Should be within mean ± jitter range.
suite.GreaterOrEqual(length, int64(noise.Mean-noise.Jitter))
suite.LessOrEqual(length, int64(noise.Mean+noise.Jitter))
}
func TestSendServerHello(t *testing.T) { func TestSendServerHello(t *testing.T) {
t.Parallel() t.Parallel()
suite.Run(t, &SendServerHelloTestSuite{}) suite.Run(t, &SendServerHelloTestSuite{})
@@ -0,0 +1,8 @@
{
"time": 1617181365,
"random": "w4TaDfYg/aUKdx1oi68vxMKvHJczRNvtRRppLETzeNE=",
"sessionId": "St2BZ2uHMFn3B2trD1jfdtpjoJOOg6JBeLhFcyCMCq4=",
"host": "storage.googleapis.com",
"cipherSuite": 4867,
"full": "FgMBAgIBAAH+AwPDhNoN9iD9pQp3HWiLry/Ewq8clzNE2+1FGmksRPN40SBK3YFna4cwWfcHa2sPWN922mOgk46DokF4uEVzIIwKrgA2WloTAxMBEwLALMArwCTAI8AKwAnMqcAwwC/AKMAnwBTAE8yoAJ0AnAA9ADwANQAvwAjAEgAKAQABf/8BAAEAAAAAGwAZAAAWc3RvcmFnZS5nb29nbGVhcGlzLmNvbQAXAAAADQAYABYEAwgEBAEFAwIDCAUIBQUBCAYGAQIBAAUABQEAAAAAM3QAAAASAAAAEAAwAC4CaDIFaDItMTYFaDItMTUFaDItMTQIc3BkeS8zLjEGc3BkeS8zCGh0dHAvMS4xAAsAAgEAADMAJgAkAB0AIAf+6C8fSRJSAC7CyUvdR9kDclNR9KLCsCFHpVZ3bC8iAC0AAgEBACsACQgDBAMDAwIDAQAKAAoACAAdABcAGAAZABUAoQAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA"
}
+158
View File
@@ -0,0 +1,158 @@
package fake
import (
"bytes"
"encoding/binary"
"errors"
"fmt"
"io"
"github.com/9seconds/mtg/v2/mtglib/internal/tls"
)
const (
maxFragmentsCount = 10
)
var ErrTooManyFragments = errors.New("too many fragments")
// https://datatracker.ietf.org/doc/html/rfc5246#section-6.2.1
// client hello can be fragmented in a series of packets:
//
// Bytes on the wire:
//
// 16 03 01 00 F8 01 00 00 F4 03 03 [32 bytes random] [session_id] [ciphers] [SNI...]
// ├─────────────┤├──────────────────────────────────────────────────────────────────┤
//
// TLS record Payload (248 bytes)
// header (5B)
//
// 16 = Handshake
// 03 01 = TLS 1.0 (record layer version)
// 00 F8 = 248 bytes follow
//
// 01 = ClientHello (handshake type)
// 00 00 F4 = 244 bytes of handshake body
// 03 03 = TLS 1.2 (actual protocol version)
// ...rest of ClientHello...
//
// Fragmented record look like:
//
// Record 1:
//
// 16 03 01 00 03 01 00 00
// ├─────────────┤├──────┤
//
// TLS header 3 bytes of payload
//
// 16 = Handshake
// 03 01 = TLS 1.0
// 00 03 = only 3 bytes follow
//
// 01 = ClientHello type
// 00 00 = first 2 bytes of the uint24 length (INCOMPLETE!)
//
// Record 2:
// 16 03 01 00 F5 F4 03 03 [32 bytes random] [session_id] [ciphers] [SNI...]
// ├─────────────┤├────────────────────────────────────────────────────────────┤
//
// TLS header remaining 245 bytes of payload
//
// 16 = Handshake
// 03 01 = TLS 1.0
// 00 F5 = 245 bytes follow
//
// F4 = last byte of uint24 length (now complete: 00 00 F4 = 244)
// 03 03 = TLS 1.2
// ...rest of ClientHello continues...
//
// So it means that there could be a series of handshake packets of different
// lengths. The goal of this function is to concatenate these fragments.
type fragmentedHandshakeReader struct {
r io.Reader
buf bytes.Buffer
readFragments int
}
func (f *fragmentedHandshakeReader) Read(p []byte) (int, error) {
if n, err := f.buf.Read(p); err == nil {
return n, nil
}
f.buf.Reset()
for f.buf.Len() == 0 {
if f.readFragments > maxFragmentsCount {
return 0, ErrTooManyFragments
}
if err := f.parseNextFragment(); err != nil {
return 0, err
}
f.readFragments++
}
return f.buf.Read(p)
}
func (f *fragmentedHandshakeReader) parseNextFragment() error {
// record_type(1) + version(2) + size(2)
// 16 - type is 0x16 (handshake record)
// 03 01 - protocol version is "3,1" (also known as TLS 1.0)
// 00 f8 - 0xF8 (248) bytes of handshake message follows
header := [1 + 2 + 2]byte{}
if _, err := io.ReadFull(f.r, header[:]); err != nil {
return fmt.Errorf("cannot read record header: %w", err)
}
if header[0] != tls.TypeHandshake {
return fmt.Errorf("unexpected record type %#x", header[0])
}
if header[1] != 3 || header[2] != 1 {
return fmt.Errorf("unexpected protocol version %#x %#x", header[1], header[2])
}
length := int64(binary.BigEndian.Uint16(header[3:]))
_, err := io.CopyN(&f.buf, f.r, length)
return err
}
func parseClientHello(r io.Reader) (*bytes.Buffer, *bytes.Buffer, error) {
r = &fragmentedHandshakeReader{r: r}
header := [1 + 3]byte{}
if _, err := io.ReadFull(r, header[:]); err != nil {
return nil, nil, fmt.Errorf("cannot read handshake header: %w", err)
}
if header[0] != TypeHandshakeClient {
return nil, nil, fmt.Errorf("incorrect handshake type: %#x", header[0])
}
// unfortunately there is not uint24 in golang, so we just reuse header
header[0] = 0
length := int64(binary.BigEndian.Uint32(header[:]))
clientHelloCopy := &bytes.Buffer{}
clientHelloCopy.Write([]byte{tls.TypeHandshake, 3, 1})
binary.Write( //nolint: errcheck
clientHelloCopy,
binary.BigEndian,
// 1 for handshake type
// 3 for handshake length
uint16(1+3+length),
)
clientHelloCopy.WriteByte(TypeHandshakeClient)
clientHelloCopy.Write(header[1:])
handshakeCopy := &bytes.Buffer{}
writer := io.MultiWriter(clientHelloCopy, handshakeCopy)
_, err := io.CopyN(writer, r, length)
return clientHelloCopy, handshakeCopy, err
}
+15 -11
View File
@@ -29,20 +29,24 @@ func ReadRecord(r io.Reader, w io.Writer) (byte, int64, error) {
func WriteRecord(w io.Writer, payload []byte) error { func WriteRecord(w io.Writer, payload []byte) error {
buf := [MaxRecordSize]byte{} buf := [MaxRecordSize]byte{}
buf[0] = TypeApplicationData copy(buf[SizeHeader:], payload)
bufV := buf[SizeRecordType:] return WriteRecordInPlace(w, buf[:], len(payload))
copy(bufV[:SizeVersion], TLSVersion[:])
bufS := bufV[SizeVersion:]
binary.BigEndian.PutUint16(bufS[:SizeSize], uint16(len(payload)))
bufP := buf[SizeHeader:]
if n := copy(bufP, payload); n != len(payload) {
return fmt.Errorf("copied %d bytes of payload instead of %d", n, len(payload))
} }
_, err := w.Write(buf[:SizeHeader+len(payload)]) func WriteRecordInPlace(w io.Writer, buf []byte, payloadLen int) error {
if payloadLen > MaxRecordPayloadSize {
return fmt.Errorf("payload %d exceeds max %d", payloadLen, MaxRecordPayloadSize)
}
buf[0] = TypeApplicationData
copy(buf[SizeRecordType:SizeRecordType+SizeVersion], TLSVersion[:])
binary.BigEndian.PutUint16(
buf[SizeRecordType+SizeVersion:SizeRecordType+SizeVersion+SizeSize],
uint16(payloadLen),
)
_, err := w.Write(buf[:SizeHeader+payloadLen])
return err return err
} }
+78
View File
@@ -119,6 +119,84 @@ func (suite *UtilsTestSuite) TestWriteRecordPayloadTooLarge() {
suite.Error(err) suite.Error(err)
} }
func (suite *UtilsTestSuite) TestWriteRecordInPlace() {
payload := []byte("hello in-place")
var buf [MaxRecordSize]byte
copy(buf[SizeHeader:], payload)
err := WriteRecordInPlace(suite.dst, buf[:], len(payload))
suite.NoError(err)
written := suite.dst.Bytes()
suite.Equal(byte(TypeApplicationData), written[0])
suite.Equal(TLSVersion[:], written[SizeRecordType:SizeRecordType+SizeVersion])
length := binary.BigEndian.Uint16(written[SizeRecordType+SizeVersion:])
suite.Equal(uint16(len(payload)), length)
suite.Equal(payload, written[SizeHeader:])
}
func (suite *UtilsTestSuite) TestWriteRecordInPlaceRoundTrip() {
payload := []byte("round trip in-place")
var buf [MaxRecordSize]byte
copy(buf[SizeHeader:], payload)
var wire bytes.Buffer
err := WriteRecordInPlace(&wire, buf[:], len(payload))
suite.NoError(err)
var recovered bytes.Buffer
recordType, length, err := ReadRecord(&wire, &recovered)
suite.NoError(err)
suite.Equal(byte(TypeApplicationData), recordType)
suite.Equal(int64(len(payload)), length)
suite.Equal(payload, recovered.Bytes())
}
func (suite *UtilsTestSuite) TestWriteRecordInPlacePayloadTooLarge() {
var buf [MaxRecordSize]byte
err := WriteRecordInPlace(suite.dst, buf[:], MaxRecordPayloadSize+1)
suite.Error(err)
}
func (suite *UtilsTestSuite) TestWriteRecordInPlacePropagatesError() {
m := &WriterMock{}
m.
On("Write", mock.AnythingOfType("[]uint8")).
Once().
Return(0, errors.New("disk full"))
var buf [MaxRecordSize]byte
copy(buf[SizeHeader:], []byte("data"))
err := WriteRecordInPlace(m, buf[:], 4)
suite.Error(err)
m.AssertExpectations(suite.T())
}
func (suite *UtilsTestSuite) TestWriteRecordInPlaceMatchesWriteRecord() {
payload := []byte("equivalence check")
var legacy bytes.Buffer
err := WriteRecord(&legacy, payload)
suite.NoError(err)
var buf [MaxRecordSize]byte
copy(buf[SizeHeader:], payload)
var inPlace bytes.Buffer
err = WriteRecordInPlace(&inPlace, buf[:], len(payload))
suite.NoError(err)
suite.Equal(legacy.Bytes(), inPlace.Bytes())
}
func TestUtils(t *testing.T) { func TestUtils(t *testing.T) {
t.Parallel() t.Parallel()
suite.Run(t, &UtilsTestSuite{}) suite.Run(t, &UtilsTestSuite{})
+38 -9
View File
@@ -27,6 +27,8 @@ type Proxy struct {
allowFallbackOnUnknownDC bool allowFallbackOnUnknownDC bool
tolerateTimeSkewness time.Duration tolerateTimeSkewness time.Duration
idleTimeout time.Duration
handshakeTimeout time.Duration
domainFrontingPort int domainFrontingPort int
domainFrontingIP string domainFrontingIP string
domainFrontingProxyProtocol bool domainFrontingProxyProtocol bool
@@ -65,10 +67,15 @@ func (p *Proxy) ServeConn(conn essentials.Conn) {
ctx := newStreamContext(p.ctx, p.logger, conn) ctx := newStreamContext(p.ctx, p.logger, conn)
defer ctx.Close() defer ctx.Close()
go func() { if err := ctx.clientConn.SetDeadline(time.Now().Add(p.handshakeTimeout)); err != nil {
<-ctx.Done() ctx.logger.WarningError("cannot set handshake timeout", err)
return
}
stop := context.AfterFunc(ctx, func() {
ctx.Close() ctx.Close()
}() })
defer stop()
p.eventStream.Send(ctx, NewEventStart(ctx.streamID, ctx.ClientIP())) p.eventStream.Send(ctx, NewEventStart(ctx.streamID, ctx.ClientIP()))
ctx.logger.Info("Stream has been started") ctx.logger.Info("Stream has been started")
@@ -96,16 +103,23 @@ func (p *Proxy) ServeConn(conn essentials.Conn) {
return return
} }
if err := ctx.clientConn.SetDeadline(time.Time{}); err != nil {
ctx.logger.WarningError("cannot set deadline", err)
return
}
if err := p.doTelegramCall(ctx); err != nil { if err := p.doTelegramCall(ctx); err != nil {
ctx.logger.WarningError("cannot dial to telegram", err) ctx.logger.WarningError("cannot dial to telegram", err)
return return
} }
tracker := newIdleTracker(p.idleTimeout)
relay.Relay( relay.Relay(
ctx, ctx,
ctx.logger.Named("relay"), ctx.logger.Named("relay"),
ctx.telegramConn, connIdleTimeout{Conn: ctx.telegramConn, tracker: tracker},
ctx.clientConn, connIdleTimeout{Conn: ctx.clientConn, tracker: tracker},
) )
} }
@@ -151,6 +165,7 @@ func (p *Proxy) Serve(listener net.Listener) error {
case errors.Is(err, ants.ErrPoolClosed): case errors.Is(err, ants.ErrPoolClosed):
return nil return nil
case errors.Is(err, ants.ErrPoolOverload): case errors.Is(err, ants.ErrPoolOverload):
conn.Close() //nolint: errcheck
logger.Info("connection was concurrency limited") logger.Info("connection was concurrency limited")
p.eventStream.Send(p.ctx, NewEventConcurrencyLimited()) p.eventStream.Send(p.ctx, NewEventConcurrencyLimited())
} }
@@ -192,7 +207,10 @@ func (p *Proxy) doFakeTLSHandshake(ctx *streamContext) bool {
return false return false
} }
if err := fake.SendServerHello(ctx.clientConn, p.secret.Key[:], clientHello); err != nil { gangerNoise := p.doppelGanger.NoiseParams()
noiseParams := fake.NoiseParams{Mean: gangerNoise.Mean, Jitter: gangerNoise.Jitter}
if err := fake.SendServerHello(ctx.clientConn, p.secret.Key[:], clientHello, noiseParams); err != nil {
p.logger.InfoError("cannot send welcome packet", err) p.logger.InfoError("cannot send welcome packet", err)
return false return false
} }
@@ -259,9 +277,16 @@ func (p *Proxy) doTelegramCall(ctx *streamContext) error {
ctx: ctx, ctx: ctx,
} }
telegramHost, _, err := net.SplitHostPort(foundAddr.Address)
if err != nil {
conn.Close() //nolint: errcheck
return fmt.Errorf("cannot parse telegram address %s: %w", foundAddr.Address, err)
}
p.eventStream.Send(ctx, p.eventStream.Send(ctx,
NewEventConnectedToDC(ctx.streamID, NewEventConnectedToDC(ctx.streamID,
conn.RemoteAddr().(*net.TCPAddr).IP, //nolint: forcetypeassert net.ParseIP(telegramHost),
ctx.dc), ctx.dc),
) )
@@ -293,11 +318,13 @@ func (p *Proxy) doDomainFronting(ctx *streamContext, conn *connRewind) {
stream: p.eventStream, stream: p.eventStream,
} }
tracker := newIdleTracker(p.idleTimeout)
relay.Relay( relay.Relay(
ctx, ctx,
ctx.logger.Named("domain-fronting"), ctx.logger.Named("domain-fronting"),
frontConn, connIdleTimeout{Conn: frontConn, tracker: tracker},
conn, connIdleTimeout{Conn: conn, tracker: tracker},
) )
} }
@@ -329,6 +356,8 @@ func NewProxy(opts ProxyOpts) (*Proxy, error) {
domainFrontingPort: opts.getDomainFrontingPort(), domainFrontingPort: opts.getDomainFrontingPort(),
domainFrontingIP: opts.DomainFrontingIP, domainFrontingIP: opts.DomainFrontingIP,
tolerateTimeSkewness: opts.getTolerateTimeSkewness(), tolerateTimeSkewness: opts.getTolerateTimeSkewness(),
idleTimeout: opts.getIdleTimeout(),
handshakeTimeout: opts.getHandshakeTimeout(),
allowFallbackOnUnknownDC: opts.AllowFallbackOnUnknownDC, allowFallbackOnUnknownDC: opts.AllowFallbackOnUnknownDC,
telegram: tg, telegram: tg,
doppelGanger: doppel.NewGanger( doppelGanger: doppel.NewGanger(
+22
View File
@@ -70,6 +70,12 @@ type ProxyOpts struct {
// This is an optional setting. // This is an optional setting.
IdleTimeout time.Duration IdleTimeout time.Duration
// HandshakeTimeout is a timeout during which all handshake ceremonies must
// be completed, otherwise this process will be aborted
//
// This is an optional setting.
HandshakeTimeout time.Duration
// TolerateTimeSkewness is a time boundary that defines a time range where // TolerateTimeSkewness is a time boundary that defines a time range where
// faketls timestamp is acceptable. // faketls timestamp is acceptable.
// //
@@ -215,6 +221,22 @@ func (p ProxyOpts) getPreferIP() string {
return p.PreferIP return p.PreferIP
} }
func (p ProxyOpts) getHandshakeTimeout() time.Duration {
if p.HandshakeTimeout == 0 {
return DefaultHandshakeTimeout
}
return p.HandshakeTimeout
}
func (p ProxyOpts) getIdleTimeout() time.Duration {
if p.IdleTimeout == 0 {
return DefaultIdleTimeout
}
return p.IdleTimeout
}
func (p ProxyOpts) getLogger(name string) Logger { func (p ProxyOpts) getLogger(name string) Logger {
return p.Logger.Named(name) return p.Logger.Named(name)
} }
+1 -1
View File
@@ -175,7 +175,7 @@ func (suite *ProxyTestSuite) TestHTTPSRequest() {
addr := fmt.Sprintf("https://%s/headers", suite.ProxyAddress()) addr := fmt.Sprintf("https://%s/headers", suite.ProxyAddress())
resp, err := client.Get(addr) //nolint: noctx resp, err := client.Get(addr) //nolint: noctx
suite.NoError(err) suite.Require().NoError(err)
defer resp.Body.Close() //nolint: errcheck defer resp.Body.Close() //nolint: errcheck
+14
View File
@@ -36,8 +36,22 @@ const (
// DefaultTCPKeepAlivePeriod defines a time period between 2 consequitive // DefaultTCPKeepAlivePeriod defines a time period between 2 consequitive
// probes. // probes.
//
// Deprecated: use DefaultKeepAliveIdle and DefaultKeepAliveInterval instead.
DefaultTCPKeepAlivePeriod = 10 * time.Second DefaultTCPKeepAlivePeriod = 10 * time.Second
// DefaultKeepAliveIdle is the time a connection must be idle before
// the first keepalive probe is sent.
DefaultKeepAliveIdle = 30 * time.Second
// DefaultKeepAliveInterval is the time between consecutive keepalive
// probes.
DefaultKeepAliveInterval = 10 * time.Second
// DefaultKeepAliveCount is the number of unacknowledged probes before
// the connection is considered dead.
DefaultKeepAliveCount = 3
// ProxyDialerOpenThreshold is used for load balancing SOCKS5 dialer only. // ProxyDialerOpenThreshold is used for load balancing SOCKS5 dialer only.
// //
// This dialer uses circuit breaker with of 3 stages: OPEN, HALF_OPEN and // This dialer uses circuit breaker with of 3 stages: OPEN, HALF_OPEN and
+7 -2
View File
@@ -20,8 +20,13 @@ func SetServerSocketOptions(conn net.Conn, bufferSize int) error {
} }
func setCommonSocketOptions(conn *net.TCPConn) error { func setCommonSocketOptions(conn *net.TCPConn) error {
if err := conn.SetKeepAlivePeriod(DefaultTCPKeepAlivePeriod); err != nil { if err := conn.SetKeepAliveConfig(net.KeepAliveConfig{
return fmt.Errorf("cannot set time period of TCP keepalive probes: %w", err) Enable: true,
Idle: DefaultKeepAliveIdle,
Interval: DefaultKeepAliveInterval,
Count: DefaultKeepAliveCount,
}); err != nil {
return fmt.Errorf("cannot configure TCP keepalive: %w", err)
} }
if err := conn.SetLinger(tcpLingerTimeout); err != nil { if err := conn.SetLinger(tcpLingerTimeout); err != nil {
+93
View File
@@ -0,0 +1,93 @@
//go:build linux || darwin
// +build linux darwin
package network_test
import (
"net"
"runtime"
"syscall"
"testing"
"time"
"github.com/9seconds/mtg/v2/network"
"github.com/stretchr/testify/require"
"golang.org/x/sys/unix"
)
func tcpKeepIdleOption() int {
if runtime.GOOS == "darwin" {
return 0x10 // TCP_KEEPALIVE on macOS
}
return 0x4 // TCP_KEEPIDLE on Linux
}
func TestSetClientSocketOptionsKeepAlive(t *testing.T) {
t.Parallel()
listener, err := net.Listen("tcp", "127.0.0.1:0")
require.NoError(t, err)
defer func() {
err := listener.Close()
require.NoError(t, err)
}()
type dialResult struct {
conn net.Conn
err error
}
dialDone := make(chan dialResult, 1)
go func() {
c, err := net.Dial("tcp", listener.Addr().String())
dialDone <- dialResult{conn: c, err: err}
}()
tcpListener, ok := listener.(*net.TCPListener)
require.True(t, ok, "listener must be a *net.TCPListener")
require.NoError(t, tcpListener.SetDeadline(time.Now().Add(5*time.Second)))
accepted, err := listener.Accept()
require.NoError(t, err)
defer func() {
err := accepted.Close()
require.NoError(t, err)
}()
dr := <-dialDone
require.NoError(t, dr.err)
defer func() {
err := dr.conn.Close()
require.NoError(t, err)
}()
err = network.SetClientSocketOptions(accepted, 0)
require.NoError(t, err)
tcpConn := accepted.(*net.TCPConn)
rawConn, err := tcpConn.SyscallConn()
require.NoError(t, err)
err = rawConn.Control(func(fd uintptr) {
val, err := unix.GetsockoptInt(int(fd), syscall.SOL_SOCKET, syscall.SO_KEEPALIVE)
require.NoError(t, err)
require.NotEqual(t, 0, val, "SO_KEEPALIVE should be enabled")
idle, err := unix.GetsockoptInt(int(fd), syscall.IPPROTO_TCP, tcpKeepIdleOption())
require.NoError(t, err)
require.Equal(t, int(network.DefaultKeepAliveIdle.Seconds()), idle)
interval, err := unix.GetsockoptInt(int(fd), syscall.IPPROTO_TCP, unix.TCP_KEEPINTVL)
require.NoError(t, err)
require.Equal(t, int(network.DefaultKeepAliveInterval.Seconds()), interval)
count, err := unix.GetsockoptInt(int(fd), syscall.IPPROTO_TCP, unix.TCP_KEEPCNT)
require.NoError(t, err)
require.Equal(t, network.DefaultKeepAliveCount, count)
})
require.NoError(t, err)
}
+1 -1
View File
@@ -25,7 +25,7 @@ func (suite *BaseHTTPTestSuite) SetupSuite() {
} }
func (suite *BaseHTTPTestSuite) SetupTest() { func (suite *BaseHTTPTestSuite) SetupTest() {
suite.client = network.New(nil, "mtg/1", 0, 0, 0).MakeHTTPClient(nil) suite.client = network.New(nil, "mtg/1", 0, 0, 0, network.DefaultKeepAliveConfig).MakeHTTPClient(nil)
} }
func (suite *BaseHTTPTestSuite) TestGet() { func (suite *BaseHTTPTestSuite) TestGet() {
+1 -1
View File
@@ -19,7 +19,7 @@ type BaseNetworkTestSuite struct {
func (suite *BaseNetworkTestSuite) SetupSuite() { func (suite *BaseNetworkTestSuite) SetupSuite() {
suite.EchoServerTestSuite.SetupSuite() suite.EchoServerTestSuite.SetupSuite()
suite.net = network.New(nil, "agent", 0, 0, 0) suite.net = network.New(nil, "agent", 0, 0, 0, network.DefaultKeepAliveConfig)
} }
func (suite *BaseNetworkTestSuite) TestDialUnknownNetwork() { func (suite *BaseNetworkTestSuite) TestDialUnknownNetwork() {
+43 -1
View File
@@ -11,6 +11,7 @@ package network
import ( import (
"errors" "errors"
"net"
"time" "time"
) )
@@ -26,14 +27,55 @@ const (
// DefaultTCPKeepAlivePeriod defines a time period between 2 consecuitive // DefaultTCPKeepAlivePeriod defines a time period between 2 consecuitive
// probes. // probes.
//
// Deprecated: use DefaultKeepAliveConfig
DefaultTCPKeepAlivePeriod = 10 * time.Second DefaultTCPKeepAlivePeriod = 10 * time.Second
// DefaultKeepAliveIdle is the time a connection must be idle before
// the first keepalive probe is sent.
//
// Deprecated: use DefaultKeepAliveConfig
DefaultKeepAliveIdle = 30 * time.Second
// DefaultKeepAliveInterval is the time between consecutive keepalive
// probes.
//
// Deprecated: use DefaultKeepAliveConfig
DefaultKeepAliveInterval = 10 * time.Second
// DefaultKeepAliveCount is the number of unacknowledged probes before
// the connection is considered dead.
//
// Deprecated: use DefaultKeepAliveConfig
DefaultKeepAliveCount = 3
// User Agent to use in HTTP client. // User Agent to use in HTTP client.
UserAgent = "curl/8.5.0" UserAgent = "curl/8.5.0"
// tcpLingerTimeout defines a number of seconds to wait for sending // tcpLingerTimeout defines a number of seconds to wait for sending
// unacknowledged data. // unacknowledged data.
tcpLingerTimeout = 1 tcpLingerTimeout = 1
// tcpNotSentLowat limits the amount of unsent data queued in the
// kernel write buffer per socket. When the unsent data drops below
// this threshold, the socket becomes writable again. This reduces
// per-connection memory usage and bufferbloat by applying
// back-pressure to the relay loop instead of piling up data in
// kernel buffers.
tcpNotSentLowat = 128 * 1024
) )
var ErrCannotDial = errors.New("cannot dial to any address") var (
ErrCannotDial = errors.New("cannot dial to any address")
// DefaultKeepAliveConfig defines a default configuration for
// keep alive settings. As per official documentation, if keep alive
// is enabled, then:
//
// Idle = 15 * time.Second
// Interval = 15 * time.Second
// Count = 9
DefaultKeepAliveConfig = net.KeepAliveConfig{
Enable: true,
}
)
+4 -1
View File
@@ -14,6 +14,7 @@ import (
type network struct { type network struct {
net.Dialer net.Dialer
keepAliveConfig net.KeepAliveConfig
httpTimeout time.Duration httpTimeout time.Duration
idleTimeout time.Duration idleTimeout time.Duration
userAgent string userAgent string
@@ -37,7 +38,7 @@ func (n *network) DialContext(ctx context.Context, network, address string) (ess
tcpConn := conn.(*net.TCPConn) tcpConn := conn.(*net.TCPConn)
return tcpConn, setCommonSocketOptions(tcpConn) return tcpConn, setCommonSocketOptions(tcpConn, n.keepAliveConfig)
} }
func (n *network) MakeHTTPClient( func (n *network) MakeHTTPClient(
@@ -71,6 +72,7 @@ func New(
tcpTimeout, tcpTimeout,
httpTimeout, httpTimeout,
idleTimeout time.Duration, idleTimeout time.Duration,
keepAliveConfig net.KeepAliveConfig,
) mtglib.Network { ) mtglib.Network {
if dnsResolver == nil { if dnsResolver == nil {
dnsResolver = net.DefaultResolver dnsResolver = net.DefaultResolver
@@ -89,5 +91,6 @@ func New(
userAgent: userAgent, userAgent: userAgent,
idleTimeout: idleTimeout, idleTimeout: idleTimeout,
httpTimeout: httpTimeout, httpTimeout: httpTimeout,
keepAliveConfig: keepAliveConfig,
} }
} }
+15
View File
@@ -3,6 +3,7 @@ package network
import ( import (
"context" "context"
"fmt" "fmt"
"net/http"
"net/url" "net/url"
"github.com/9seconds/mtg/v2/essentials" "github.com/9seconds/mtg/v2/essentials"
@@ -15,6 +16,10 @@ type proxyNetwork struct {
client proxy.ContextDialer client proxy.ContextDialer
} }
func (p proxyNetwork) Dial(network, address string) (essentials.Conn, error) {
return p.DialContext(context.Background(), network, address)
}
func (p proxyNetwork) DialContext(ctx context.Context, network, address string) (essentials.Conn, error) { func (p proxyNetwork) DialContext(ctx context.Context, network, address string) (essentials.Conn, error) {
conn, err := p.client.DialContext(ctx, network, address) conn, err := p.client.DialContext(ctx, network, address)
if err != nil { if err != nil {
@@ -24,6 +29,16 @@ func (p proxyNetwork) DialContext(ctx context.Context, network, address string)
return essentials.WrapNetConn(conn), nil return essentials.WrapNetConn(conn), nil
} }
func (p proxyNetwork) MakeHTTPClient(
dialFunc func(context.Context, string, string) (essentials.Conn, error),
) *http.Client {
if dialFunc == nil {
dialFunc = p.DialContext
}
return p.Network.MakeHTTPClient(dialFunc)
}
func NewProxyNetwork(base mtglib.Network, proxyURL *url.URL) (*proxyNetwork, error) { func NewProxyNetwork(base mtglib.Network, proxyURL *url.URL) (*proxyNetwork, error) {
socks, err := proxy.FromURL(proxyURL, base.NativeDialer()) socks, err := proxy.FromURL(proxyURL, base.NativeDialer())
if err != nil { if err != nil {
+7 -3
View File
@@ -5,9 +5,9 @@ import (
"net" "net"
) )
func setCommonSocketOptions(conn *net.TCPConn) error { func setCommonSocketOptions(conn *net.TCPConn, keepAliveConfig net.KeepAliveConfig) error {
if err := conn.SetKeepAlivePeriod(DefaultTCPKeepAlivePeriod); err != nil { if err := conn.SetKeepAliveConfig(keepAliveConfig); err != nil {
return fmt.Errorf("cannot set time period of TCP keepalive probes: %w", err) return fmt.Errorf("cannot configure TCP keepalive: %w", err)
} }
if err := conn.SetLinger(tcpLingerTimeout); err != nil { if err := conn.SetLinger(tcpLingerTimeout); err != nil {
@@ -23,5 +23,9 @@ func setCommonSocketOptions(conn *net.TCPConn) error {
return fmt.Errorf("cannot setup SO_REUSEADDR/PORT: %w", err) return fmt.Errorf("cannot setup SO_REUSEADDR/PORT: %w", err)
} }
setCongestionControl(rawConn)
setTCPUserTimeout(rawConn, keepAliveConfig)
setNotSentLowat(rawConn)
return nil return nil
} }
+20
View File
@@ -0,0 +1,20 @@
//go:build linux
package network
import (
"syscall"
"golang.org/x/sys/unix"
)
// setCongestionControl sets BBR as the TCP congestion control algorithm.
// BBR provides better throughput over lossy and high-latency links compared
// to the default cubic, which is especially beneficial for mobile and
// home internet clients. This is best-effort: silently ignored if the
// kernel does not have tcp_bbr available.
func setCongestionControl(conn syscall.RawConn) {
conn.Control(func(fd uintptr) { //nolint: errcheck
unix.SetsockoptString(int(fd), unix.IPPROTO_TCP, unix.TCP_CONGESTION, "bbr") //nolint: errcheck
})
}
+7
View File
@@ -0,0 +1,7 @@
//go:build !linux
package network
import "syscall"
func setCongestionControl(conn syscall.RawConn) {}
+20
View File
@@ -0,0 +1,20 @@
//go:build linux || darwin
package network
import (
"syscall"
"golang.org/x/sys/unix"
)
// setNotSentLowat sets TCP_NOTSENT_LOWAT which limits the amount of
// unsent data queued in the kernel write buffer. Once unsent data drops
// below this threshold the socket becomes writable again, applying
// back-pressure to the relay loop instead of piling up data in kernel
// buffers. This reduces per-connection memory and bufferbloat.
func setNotSentLowat(conn syscall.RawConn) {
conn.Control(func(fd uintptr) { //nolint: errcheck
unix.SetsockoptInt(int(fd), unix.IPPROTO_TCP, unix.TCP_NOTSENT_LOWAT, tcpNotSentLowat) //nolint: errcheck
})
}
+7
View File
@@ -0,0 +1,7 @@
//go:build !linux && !darwin
package network
import "syscall"
func setNotSentLowat(conn syscall.RawConn) {}
@@ -1,5 +1,4 @@
//go:build !windows //go:build !windows
// +build !windows
package network package network
@@ -1,5 +1,4 @@
//go:build windows //go:build windows
// +build windows
package network package network
+92
View File
@@ -0,0 +1,92 @@
//go:build linux || darwin
// +build linux darwin
package network
import (
"net"
"runtime"
"syscall"
"testing"
"time"
"github.com/stretchr/testify/require"
"golang.org/x/sys/unix"
)
func tcpKeepIdleOption() int {
if runtime.GOOS == "darwin" {
return 0x10 // TCP_KEEPALIVE on macOS
}
return 0x4 // TCP_KEEPIDLE on Linux
}
func TestSetCommonSocketOptionsKeepAlive(t *testing.T) {
t.Parallel()
listener, err := net.Listen("tcp", "127.0.0.1:0")
require.NoError(t, err)
defer func() {
err := listener.Close()
require.NoError(t, err)
}()
type dialResult struct {
conn net.Conn
err error
}
dialDone := make(chan dialResult, 1)
go func() {
c, err := net.Dial("tcp", listener.Addr().String())
dialDone <- dialResult{conn: c, err: err}
}()
tcpListener, ok := listener.(*net.TCPListener)
require.True(t, ok, "listener must be a *net.TCPListener")
require.NoError(t, tcpListener.SetDeadline(time.Now().Add(5*time.Second)))
accepted, err := listener.Accept()
require.NoError(t, err)
defer func() {
err := accepted.Close()
require.NoError(t, err)
}()
dr := <-dialDone
require.NoError(t, dr.err)
defer func() {
err := dr.conn.Close()
require.NoError(t, err)
}()
tcpConn := accepted.(*net.TCPConn)
err = setCommonSocketOptions(tcpConn, DefaultKeepAliveConfig)
require.NoError(t, err)
rawConn, err := tcpConn.SyscallConn()
require.NoError(t, err)
err = rawConn.Control(func(fd uintptr) {
val, err := unix.GetsockoptInt(int(fd), syscall.SOL_SOCKET, syscall.SO_KEEPALIVE)
require.NoError(t, err)
require.NotEqual(t, 0, val, "SO_KEEPALIVE should be enabled")
idle, err := unix.GetsockoptInt(int(fd), syscall.IPPROTO_TCP, tcpKeepIdleOption())
require.NoError(t, err)
require.Equal(t, 15, idle, "keepalive idle should match DefaultKeepAliveIdle")
interval, err := unix.GetsockoptInt(int(fd), syscall.IPPROTO_TCP, unix.TCP_KEEPINTVL)
require.NoError(t, err)
require.Equal(t, 15, interval, "keepalive interval should match DefaultKeepAliveInterval")
count, err := unix.GetsockoptInt(int(fd), syscall.IPPROTO_TCP, unix.TCP_KEEPCNT)
require.NoError(t, err)
require.Equal(t, 9, count, "keepalive count should match DefaultKeepAliveCount")
})
require.NoError(t, err)
}
+48
View File
@@ -0,0 +1,48 @@
//go:build linux
package network
import (
"net"
"syscall"
"time"
"golang.org/x/sys/unix"
)
// Go runtime defaults for KeepAliveConfig when fields are zero.
const (
goDefaultKeepAliveIdle = 15 * time.Second
goDefaultKeepAliveInterval = 15 * time.Second
goDefaultKeepAliveCount = 9
)
// setTCPUserTimeout sets TCP_USER_TIMEOUT on a socket. If transmitted
// data remains unacknowledged for this long, the kernel closes the
// connection. As recommended by Cloudflare
// (https://blog.cloudflare.com/when-tcp-sockets-refuse-to-die/),
// the value is computed as: keepidle + keepintvl * keepcnt. This
// ensures TCP_USER_TIMEOUT and keepalives agree on when to give up.
// Best-effort: silently ignored if unsupported.
func setTCPUserTimeout(conn syscall.RawConn, cfg net.KeepAliveConfig) {
idle := cfg.Idle
if idle == 0 {
idle = goDefaultKeepAliveIdle
}
interval := cfg.Interval
if interval == 0 {
interval = goDefaultKeepAliveInterval
}
count := cfg.Count
if count == 0 {
count = goDefaultKeepAliveCount
}
timeout := idle + interval*time.Duration(count)
conn.Control(func(fd uintptr) { //nolint: errcheck
unix.SetsockoptInt(int(fd), unix.IPPROTO_TCP, unix.TCP_USER_TIMEOUT, int(timeout.Milliseconds())) //nolint: errcheck
})
}
+10
View File
@@ -0,0 +1,10 @@
//go:build !linux
package network
import (
"net"
"syscall"
)
func setTCPUserTimeout(conn syscall.RawConn, cfg net.KeepAliveConfig) {}
+1 -1
View File
@@ -66,7 +66,7 @@ func (suite *SocksProxyTestSuite) SetupSuite() {
require.NoError(suite.T(), err) require.NoError(suite.T(), err)
suite.authURL = parsed suite.authURL = parsed
suite.baseNetwork = network.New(nil, "mtg", 0, 0, 0) suite.baseNetwork = network.New(nil, "mtg", 0, 0, 0, network.DefaultKeepAliveConfig)
} }
func (suite *SocksProxyTestSuite) TestIncorrectSchema() { func (suite *SocksProxyTestSuite) TestIncorrectSchema() {
+6
View File
@@ -0,0 +1,6 @@
//go:build !prof
package main
func runProfile() {
}
+26
View File
@@ -0,0 +1,26 @@
//go:build prof
package main
import (
"net"
"net/http"
_ "net/http/pprof" //nolint: gosec
"os"
)
const DefaultProfPort = "6000"
func runProfile() {
port := os.Getenv("MTG_PROF_PORT")
if port == "" {
port = DefaultProfPort
}
listener, err := net.Listen("tcp", net.JoinHostPort("127.0.0.1", port))
if err != nil {
panic(err)
}
go http.Serve(listener, nil)
}