feat(loadtest): sync add real 500-session capacity harness

This commit is contained in:
iamxvbaba 2026-08-02 12:02:07 +08:00
parent ac0566f779
commit 141f2f20c4
39 changed files with 4157 additions and 42 deletions

View file

@ -1,6 +1,7 @@
package tdesktop
import (
"net/netip"
"time"
"github.com/iamxvbaba/td/tg"
@ -13,18 +14,27 @@ import (
// 字段值取 Telegram 常见默认TDesktop 联调阶段按客户端实际需要微调
// (记录于 docs/compatibility-matrix.md
func BuildConfig(dc int, ip string, port int, now time.Time, publicBaseURL string) *tg.Config {
// TELESRV_ADVERTISE_IP is validated during config loading. Parse again here
// only to derive the wire ipv6 flag and to render IPv4-mapped addresses in
// their canonical form. Keeping the advertised route in help.getConfig is a
// protocol invariant: clients replace or persist this list for reconnects.
addr, err := netip.ParseAddr(ip)
if err == nil {
addr = addr.Unmap()
ip = addr.String()
}
meURLPrefix := links.NormalizeBaseURL(publicBaseURL) + "/"
config := &tg.Config{
Date: int(now.Unix()),
Expires: int(now.Add(time.Hour).Unix()),
TestMode: false,
ThisDC: dc,
// 不下发 DCOptions客户端TDesktop patch / drklo fork已写死 static DC
// 地址空列表会让客户端保留它——drklo ConnectionsManager.cpp 的 processConfig
// 在 dc_options 为空时整段跳过 replaceAddresses/saveConfig既不覆盖也不持久化。
// 服务端因此无需配置对外可达 IP换网络/部署只改客户端写死地址即可。ip/port
// 参数暂留,供未来需要显式 advertise 时改回。
DCOptions: nil,
DCOptions: []tg.DCOption{{
Ipv6: addr.Is6(),
ID: dc,
IPAddress: ip,
Port: port,
}},
ChatSizeMax: 200,
MegagroupSizeMax: 200000,
ForwardedCountMax: 100,

View file

@ -18,3 +18,31 @@ func TestBuildConfigIncludesDefaultReaction(t *testing.T) {
t.Fatalf("reactions_default = %#v, want %q emoji", reaction, DefaultReactionEmoticon)
}
}
func TestBuildConfigAdvertisesCanonicalPrimaryDC(t *testing.T) {
tests := []struct {
name string
ip string
want string
ipv6 bool
}{
{name: "ipv4", ip: "192.0.2.10", want: "192.0.2.10"},
{name: "ipv6", ip: "2001:0db8::1", want: "2001:db8::1", ipv6: true},
{name: "mapped ipv4", ip: "::ffff:192.0.2.10", want: "192.0.2.10"},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
config := BuildConfig(2, tt.ip, 2398, time.Unix(1, 0), "https://telesrv.net")
if len(config.DCOptions) != 1 {
t.Fatalf("len(DCOptions) = %d, want 1", len(config.DCOptions))
}
option := config.DCOptions[0]
if option.ID != 2 || option.IPAddress != tt.want || option.Port != 2398 || option.Ipv6 != tt.ipv6 {
t.Fatalf("DCOptions[0] = %+v, want dc=2 ip=%q port=2398 ipv6=%v", option, tt.want, tt.ipv6)
}
if option.MediaOnly || option.CDN || option.TCPObfuscatedOnly || option.Static || option.ThisPortOnly {
t.Fatalf("DCOptions[0] has unexpected restrictive flags: %+v", option)
}
})
}
}

View file

@ -597,6 +597,10 @@ func Load() (Config, error) {
if err != nil {
return Config{}, fmt.Errorf("TELESRV_DEFAULT_COUNTRY_CODE: %w", err)
}
advertiseIP, err := normalizeAdvertiseIP(envOr("TELESRV_ADVERTISE_IP", "127.0.0.1"))
if err != nil {
return Config{}, fmt.Errorf("TELESRV_ADVERTISE_IP: %w", err)
}
// The composite rating weight defaults are the domain formula's own defaults;
// see RatingWeight* below.
defaultRatingWeights := domain.DefaultAccountRatingWeights()
@ -612,10 +616,9 @@ func Load() (Config, error) {
"http://localhost:1234",
"http://127.0.0.1:1234",
}),
// AdvertiseIP 当前不影响 help.getConfig——getConfig 返回空 DCOptions
// 客户端使用其写死的 static DC 地址(见 compat/tdesktop/config.go
// 字段与默认值保留,供未来需要显式下发 DC 地址时使用。
AdvertiseIP: envOr("TELESRV_ADVERTISE_IP", "127.0.0.1"),
// help.getConfig 必须下发至少一个可重连的主 DC 地址;远端部署不能
// 沿用 loopback 默认值,需显式设置客户端实际可达的 IP。
AdvertiseIP: advertiseIP,
RSAKeyPath: envOr("TELESRV_RSA_KEY", "data/server_rsa.pem"),
DC: envIntOr("TELESRV_DC", 2),
DefaultCountryCode: countryCode,
@ -896,6 +899,18 @@ func normalizeDefaultCountryCode(raw string) (string, error) {
return region.String(), nil
}
func normalizeAdvertiseIP(raw string) (string, error) {
addr, err := netip.ParseAddr(strings.TrimSpace(raw))
if err != nil {
return "", fmt.Errorf("must be an IPv4 or IPv6 address: %w", err)
}
addr = addr.Unmap()
if addr.IsUnspecified() || addr.IsMulticast() || addr.Zone() != "" {
return "", fmt.Errorf("must be a unicast address usable by clients")
}
return addr.String(), nil
}
func validateTelegramLoginConfig(cfg Config) error {
if !cfg.TelegramLoginEnabled {
return nil

View file

@ -70,6 +70,32 @@ func TestLoadUsesExplicitAdvertiseIP(t *testing.T) {
}
}
func TestLoadCanonicalizesAdvertiseIP(t *testing.T) {
disableDefaultConfigFile(t)
t.Setenv("TELESRV_ADVERTISE_IP", " 2001:0db8::1 ")
cfg, err := Load()
if err != nil {
t.Fatalf("Load: %v", err)
}
if cfg.AdvertiseIP != "2001:db8::1" {
t.Fatalf("AdvertiseIP = %q, want canonical IPv6", cfg.AdvertiseIP)
}
}
func TestLoadRejectsUnusableAdvertiseIP(t *testing.T) {
for _, value := range []string{"example.com", "0.0.0.0", "::", "224.0.0.1", "fe80::1%eth0"} {
t.Run(value, func(t *testing.T) {
disableDefaultConfigFile(t)
t.Setenv("TELESRV_ADVERTISE_IP", value)
if _, err := Load(); err == nil {
t.Fatalf("Load accepted TELESRV_ADVERTISE_IP=%q", value)
}
})
}
}
func TestLoadDefaultCountryCode(t *testing.T) {
t.Run("default", func(t *testing.T) {
disableDefaultConfigFile(t)

View file

@ -0,0 +1,119 @@
package loadharness
import (
"context"
"crypto/rsa"
"crypto/x509"
"encoding/pem"
"errors"
"fmt"
"net"
"os"
"path/filepath"
"strconv"
"strings"
"github.com/iamxvbaba/td/exchange"
"github.com/iamxvbaba/td/telegram"
"github.com/iamxvbaba/td/telegram/dcs"
"github.com/iamxvbaba/td/tg"
"github.com/iamxvbaba/td/transport"
)
type clientHooks struct {
Update telegram.UpdateHandler
ConnectionState func(telegram.ConnectionState)
Dead func(error)
}
func newClient(endpoint Endpoint, publicKey *rsa.PublicKey, storage telegram.SessionStorage, hooks clientHooks) (*telegram.Client, error) {
host, portText, err := net.SplitHostPort(endpoint.Address)
if err != nil {
return nil, fmt.Errorf("parse endpoint address: %w", err)
}
port, err := strconv.Atoi(portText)
if err != nil || port <= 0 || port > 65535 {
return nil, fmt.Errorf("invalid endpoint port %q", portText)
}
protocol := dcs.Protocol(transport.Intermediate)
if endpoint.Obfuscated {
protocol = transport.Abridged
}
resolver := dcs.Plain(dcs.PlainOptions{Protocol: protocol, Obfuscated: endpoint.Obfuscated})
updateHandler := hooks.Update
if updateHandler == nil {
updateHandler = telegram.UpdateHandlerFunc(func(context.Context, tg.UpdatesClass) error { return nil })
}
return telegram.NewClient(endpoint.APIID, endpoint.APIHash, telegram.Options{
PublicKeys: []exchange.PublicKey{{RSA: publicKey}},
DC: endpoint.DC,
Resolver: resolver,
DCList: dcs.List{Options: []tg.DCOption{{
ID: endpoint.DC, IPAddress: host, Port: port, Static: true,
}}},
SessionStorage: storage,
UpdateHandler: updateHandler,
EnablePFS: endpoint.PFS,
TempKeyTTL: endpoint.TempKeyTTL,
Device: telegram.DeviceTDesktopWindows(),
OnConnectionState: hooks.ConnectionState,
OnDead: hooks.Dead,
}), nil
}
func loadRSAPublicKey(path string) (*rsa.PublicKey, error) {
data, err := os.ReadFile(path)
if err != nil {
return nil, fmt.Errorf("read RSA key: %w", err)
}
block, _ := pem.Decode(data)
if block == nil {
return nil, errors.New("RSA key is not PEM")
}
if private, err := x509.ParsePKCS1PrivateKey(block.Bytes); err == nil {
return &private.PublicKey, nil
}
if parsed, err := x509.ParsePKCS8PrivateKey(block.Bytes); err == nil {
if private, ok := parsed.(*rsa.PrivateKey); ok {
return &private.PublicKey, nil
}
}
if public, err := x509.ParsePKCS1PublicKey(block.Bytes); err == nil {
return public, nil
}
if parsed, err := x509.ParsePKIXPublicKey(block.Bytes); err == nil {
if public, ok := parsed.(*rsa.PublicKey); ok {
return public, nil
}
}
return nil, errors.New("PEM does not contain an RSA private or public key")
}
func writePortablePublicKey(manifestPath, sourcePath string) (string, *rsa.PublicKey, error) {
publicKey, err := loadRSAPublicKey(sourcePath)
if err != nil {
return "", nil, err
}
encoded, err := x509.MarshalPKIXPublicKey(publicKey)
if err != nil {
return "", nil, err
}
const name = "server_rsa_public.pem"
path := filepath.Join(filepath.Dir(manifestPath), name)
data := pem.EncodeToMemory(&pem.Block{Type: "PUBLIC KEY", Bytes: encoded})
if err := writeFileAtomic(path, data, 0o644); err != nil {
return "", nil, err
}
return name, publicKey, nil
}
func loadManifestPublicKey(manifestPath string, endpoint Endpoint, override string) (*rsa.PublicKey, error) {
path := strings.TrimSpace(override)
if path == "" {
path = endpoint.RSAKeyPath
if !filepath.IsAbs(path) {
path = filepath.Join(filepath.Dir(manifestPath), filepath.FromSlash(path))
}
}
return loadRSAPublicKey(path)
}

View file

@ -0,0 +1,112 @@
package loadharness
import (
"encoding/json"
"errors"
"fmt"
"os"
"path/filepath"
"strings"
"time"
"github.com/iamxvbaba/td/tg"
)
const (
fileFixtureVersion = 1
fixturePatternVersion = 1
)
// persistedFileFixture keeps only the stable location of a synthetic load-test
// document. It contains no auth key or login secret and is owner-readable so a
// test bundle can reuse the same server-side file across independent runs.
type persistedFileFixture struct {
Version int `json:"version"`
CreatedAt time.Time `json:"created_at"`
ServerAddress string `json:"server_address"`
DC int `json:"dc"`
SizeBytes int `json:"size_bytes"`
PatternVersion int `json:"pattern_version"`
DocumentID int64 `json:"document_id"`
AccessHash int64 `json:"access_hash"`
FileReference []byte `json:"file_reference"`
}
func (f *persistedFileFixture) validate(endpoint Endpoint, size int) error {
if f == nil {
return errors.New("nil file fixture")
}
if f.Version != fileFixtureVersion || f.PatternVersion != fixturePatternVersion {
return errors.New("file fixture version does not match the harness")
}
if f.ServerAddress != endpoint.Address || f.DC != endpoint.DC {
return errors.New("file fixture endpoint does not match the manifest")
}
if f.SizeBytes != size || f.SizeBytes <= 0 {
return fmt.Errorf("file fixture size %d does not match requested %d", f.SizeBytes, size)
}
if f.DocumentID == 0 || f.AccessHash == 0 || len(f.FileReference) == 0 {
return errors.New("file fixture has an incomplete document location")
}
return nil
}
func (f *persistedFileFixture) runtime(chunk int) *downloadFixture {
return &downloadFixture{
location: &tg.InputDocumentFileLocation{
ID: f.DocumentID, AccessHash: f.AccessHash,
FileReference: append([]byte(nil), f.FileReference...),
},
size: f.SizeBytes, chunk: chunk,
}
}
func persistedFixture(endpoint Endpoint, fixture *downloadFixture) *persistedFileFixture {
return &persistedFileFixture{
Version: fileFixtureVersion, CreatedAt: time.Now().UTC(),
ServerAddress: endpoint.Address, DC: endpoint.DC,
SizeBytes: fixture.size, PatternVersion: fixturePatternVersion,
DocumentID: fixture.location.ID, AccessHash: fixture.location.AccessHash,
FileReference: append([]byte(nil), fixture.location.FileReference...),
}
}
func resolveFileFixturePath(manifestPath, configured string) string {
configured = strings.TrimSpace(configured)
if configured == "" {
return filepath.Join(filepath.Dir(manifestPath), "file-fixture.json")
}
if filepath.IsAbs(configured) {
return configured
}
return filepath.Join(filepath.Dir(manifestPath), filepath.FromSlash(configured))
}
func loadPersistedFileFixture(path string, endpoint Endpoint, size, chunk int) (*downloadFixture, error) {
data, err := os.ReadFile(path)
if err != nil {
return nil, err
}
var fixture persistedFileFixture
decoder := json.NewDecoder(strings.NewReader(string(data)))
decoder.DisallowUnknownFields()
if err := decoder.Decode(&fixture); err != nil {
return nil, fmt.Errorf("decode file fixture: %w", err)
}
if err := fixture.validate(endpoint, size); err != nil {
return nil, err
}
return fixture.runtime(chunk), nil
}
func writePersistedFileFixture(path string, endpoint Endpoint, fixture *downloadFixture) error {
persisted := persistedFixture(endpoint, fixture)
if err := persisted.validate(endpoint, fixture.size); err != nil {
return err
}
data, err := json.MarshalIndent(persisted, "", " ")
if err != nil {
return fmt.Errorf("encode file fixture: %w", err)
}
return writeFileAtomic(path, append(data, '\n'), 0o600)
}

View file

@ -0,0 +1,44 @@
package loadharness
import (
"path/filepath"
"testing"
"github.com/iamxvbaba/td/tg"
)
func TestPersistedFileFixtureRoundTripAndIdentityChecks(t *testing.T) {
dir := t.TempDir()
path := filepath.Join(dir, "fixture.json")
endpoint := Endpoint{Address: "127.0.0.1:2398", DC: 2}
want := &downloadFixture{
location: &tg.InputDocumentFileLocation{ID: 42, AccessHash: 99, FileReference: []byte{1, 2, 3}},
size: 4 << 20, chunk: 1 << 20,
}
if err := writePersistedFileFixture(path, endpoint, want); err != nil {
t.Fatal(err)
}
got, err := loadPersistedFileFixture(path, endpoint, want.size, want.chunk)
if err != nil {
t.Fatal(err)
}
if got.size != want.size || got.chunk != want.chunk || got.location.ID != want.location.ID || got.location.AccessHash != want.location.AccessHash || string(got.location.FileReference) != string(want.location.FileReference) {
t.Fatalf("fixture = %#v, want %#v", got, want)
}
if _, err := loadPersistedFileFixture(path, Endpoint{Address: "other:2398", DC: 2}, want.size, want.chunk); err == nil {
t.Fatal("expected endpoint mismatch")
}
if _, err := loadPersistedFileFixture(path, endpoint, want.size/2, want.chunk); err == nil {
t.Fatal("expected size mismatch")
}
}
func TestResolveFileFixturePathDefaultsBesideManifest(t *testing.T) {
manifest := filepath.Join(t.TempDir(), "bundle", "manifest.json")
if got, want := resolveFileFixturePath(manifest, ""), filepath.Join(filepath.Dir(manifest), "file-fixture.json"); got != want {
t.Fatalf("default path = %q, want %q", got, want)
}
if got, want := resolveFileFixturePath(manifest, "custom.json"), filepath.Join(filepath.Dir(manifest), "custom.json"); got != want {
t.Fatalf("relative path = %q, want %q", got, want)
}
}

View file

@ -0,0 +1,5 @@
//go:build !darwin && !linux
package loadharness
func validateProcessCapacity(int) error { return nil }

View file

@ -0,0 +1,12 @@
package loadharness
import "testing"
func TestMinimumOpenFilesHasFixedAndPerSessionHeadroom(t *testing.T) {
if got, want := minimumOpenFiles(0), 256; got != want {
t.Fatalf("minimumOpenFiles(0) = %d, want %d", got, want)
}
if got, want := minimumOpenFiles(500), 3256; got != want {
t.Fatalf("minimumOpenFiles(500) = %d, want %d", got, want)
}
}

View file

@ -0,0 +1,21 @@
//go:build darwin || linux
package loadharness
import (
"fmt"
"golang.org/x/sys/unix"
)
func validateProcessCapacity(sessions int) error {
var limit unix.Rlimit
if err := unix.Getrlimit(unix.RLIMIT_NOFILE, &limit); err != nil {
return fmt.Errorf("read open-file limit: %w", err)
}
required := minimumOpenFiles(sessions)
if limit.Cur < uint64(required) {
return fmt.Errorf("open-file soft limit %d is below required %d for %d sessions; raise it before running the load", limit.Cur, required, sessions)
}
return nil
}

View file

@ -0,0 +1,267 @@
package loadharness
import (
"context"
"crypto/rsa"
"errors"
"fmt"
"os"
"path/filepath"
"sort"
"strings"
"sync"
"time"
"github.com/iamxvbaba/td/session"
"github.com/iamxvbaba/td/tg"
)
type ProvisionConfig struct {
ManifestPath string
SessionKeyPath string
RSAKeyPath string
Endpoint Endpoint
Accounts int
ExtraDevices int
Concurrency int
PhonePrefix string
Code string
FirstNamePrefix string
}
type ProvisionEvent struct {
Completed int
Total int
Session SessionRecord
Resumed bool
Err error
}
func (c ProvisionConfig) validate() error {
if err := c.Endpoint.Validate(); err != nil {
return err
}
if c.ManifestPath == "" || c.SessionKeyPath == "" || c.RSAKeyPath == "" {
return errors.New("manifest, session-key and RSA key paths are required")
}
if c.Accounts <= 0 || c.ExtraDevices < 0 || c.ExtraDevices > c.Accounts {
return errors.New("accounts must be positive and extra-devices must be between zero and accounts")
}
if c.Concurrency <= 0 || c.Concurrency > 64 {
return errors.New("provision concurrency must be between 1 and 64")
}
if strings.TrimSpace(c.Code) == "" {
return errors.New("a test login code is required")
}
return nil
}
// Provision creates accounts only through auth.sendCode/signIn/signUp. Primary
// devices finish before duplicate-device login starts, preventing two workers
// from racing the first signup for one phone.
func Provision(ctx context.Context, cfg ProvisionConfig, progress func(ProvisionEvent)) (*Manifest, error) {
if err := cfg.validate(); err != nil {
return nil, err
}
key, err := LoadSessionKey(cfg.SessionKeyPath)
if err != nil {
return nil, err
}
publicName, publicKey, err := writePortablePublicKey(cfg.ManifestPath, cfg.RSAKeyPath)
if err != nil {
return nil, err
}
cfg.Endpoint.RSAKeyPath = publicName
primary := make([]SessionRecord, 0, cfg.Accounts)
for account := 0; account < cfg.Accounts; account++ {
primary = append(primary, desiredSessionRecord(account, account, 0, cfg))
}
completed, err := provisionPhase(ctx, cfg, key, publicKey, primary, progress, 0, cfg.Accounts+cfg.ExtraDevices)
if err != nil {
return nil, err
}
extra := make([]SessionRecord, 0, cfg.ExtraDevices)
for account := 0; account < cfg.ExtraDevices; account++ {
extra = append(extra, desiredSessionRecord(cfg.Accounts+account, account, 1, cfg))
}
extraCompleted, err := provisionPhase(ctx, cfg, key, publicKey, extra, progress, len(completed), cfg.Accounts+cfg.ExtraDevices)
if err != nil {
return nil, err
}
completed = append(completed, extraCompleted...)
sort.Slice(completed, func(i, j int) bool { return completed[i].Index < completed[j].Index })
manifest := &Manifest{
Version: ManifestVersion, CreatedAt: time.Now().UTC(), Endpoint: cfg.Endpoint, Sessions: completed,
}
if err := WriteManifest(cfg.ManifestPath, manifest); err != nil {
return nil, err
}
return manifest, nil
}
func desiredSessionRecord(index, account, device int, cfg ProvisionConfig) SessionRecord {
return SessionRecord{
Index: index, AccountIndex: account, DeviceIndex: device,
Phone: fmt.Sprintf("%s%06d", cfg.PhonePrefix, account+1),
FirstName: fmt.Sprintf("%s%04d", cfg.FirstNamePrefix, account+1),
SessionFile: filepath.ToSlash(filepath.Join(sessionDirectoryForManifest(cfg.ManifestPath), fmt.Sprintf("session-%04d-device-%d.bin", account, device))),
}
}
// sessionDirectoryForManifest keeps independently named manifests in the same
// parent directory from ever sharing encrypted session files. The conventional
// manifest.json path retains the compact "sessions" directory, so moving a
// complete bundle to another host remains portable.
func sessionDirectoryForManifest(manifestPath string) string {
base := filepath.Base(filepath.Clean(manifestPath))
base = strings.TrimSuffix(base, filepath.Ext(base))
if base == "" || base == "." || strings.EqualFold(base, "manifest") {
return "sessions"
}
return "sessions-" + base
}
func provisionPhase(
ctx context.Context,
cfg ProvisionConfig,
key [32]byte,
publicKey *rsa.PublicKey,
desired []SessionRecord,
progress func(ProvisionEvent),
completedBefore, total int,
) ([]SessionRecord, error) {
if len(desired) == 0 {
return nil, nil
}
type result struct {
record SessionRecord
resumed bool
err error
}
jobs := make(chan SessionRecord)
results := make(chan result, len(desired))
workers := min(cfg.Concurrency, len(desired))
var wg sync.WaitGroup
for i := 0; i < workers; i++ {
wg.Add(1)
go func() {
defer wg.Done()
for record := range jobs {
path := resolveSessionPath(cfg.ManifestPath, record)
_, statErr := os.Stat(path)
resumed := statErr == nil
storage := &EncryptedFileStorage{Path: path, Key: key}
user, err := provisionOne(ctx, cfg, publicKey, storage, record)
if err == nil {
record.UserID = user.ID
record.AccessHash = user.AccessHash
}
results <- result{record: record, resumed: resumed, err: err}
}
}()
}
go func() {
defer close(jobs)
for _, record := range desired {
select {
case jobs <- record:
case <-ctx.Done():
return
}
}
}()
go func() { wg.Wait(); close(results) }()
completed := make([]SessionRecord, 0, len(desired))
var firstErr error
for result := range results {
if result.err == nil {
completed = append(completed, result.record)
} else if firstErr == nil {
firstErr = fmt.Errorf("provision session %d: %w", result.record.Index, result.err)
}
if progress != nil {
progress(ProvisionEvent{
Completed: completedBefore + len(completed), Total: total,
Session: result.record, Resumed: result.resumed, Err: result.err,
})
}
}
if firstErr != nil {
return nil, firstErr
}
if len(completed) != len(desired) {
return nil, ctx.Err()
}
return completed, nil
}
func provisionOne(ctx context.Context, cfg ProvisionConfig, publicKey *rsa.PublicKey, storage *EncryptedFileStorage, record SessionRecord) (*tg.User, error) {
client, err := newClient(cfg.Endpoint, publicKey, storage, clientHooks{})
if err != nil {
return nil, err
}
var user *tg.User
err = client.Run(ctx, func(ctx context.Context) error {
status, err := client.Auth().Status(ctx)
if err != nil {
return fmt.Errorf("authorization status: %w", err)
}
if status.Authorized && status.User != nil {
user = status.User
return nil
}
raw := tg.NewClient(client)
sent, err := raw.AuthSendCode(ctx, &tg.AuthSendCodeRequest{
PhoneNumber: record.Phone, APIID: cfg.Endpoint.APIID, APIHash: cfg.Endpoint.APIHash, Settings: tg.CodeSettings{},
})
if err != nil {
return fmt.Errorf("auth.sendCode: %w", err)
}
sentCode, ok := sent.(*tg.AuthSentCode)
if !ok {
return fmt.Errorf("auth.sendCode returned %T", sent)
}
authorization, err := raw.AuthSignIn(ctx, &tg.AuthSignInRequest{
PhoneNumber: record.Phone, PhoneCodeHash: sentCode.PhoneCodeHash, PhoneCode: cfg.Code,
})
if err != nil {
return fmt.Errorf("auth.signIn: %w", err)
}
if authorized, ok := authorization.(*tg.AuthAuthorization); ok {
user, ok = authorized.User.(*tg.User)
if !ok {
return fmt.Errorf("auth.signIn user is %T", authorized.User)
}
return nil
}
if _, ok := authorization.(*tg.AuthAuthorizationSignUpRequired); !ok {
return fmt.Errorf("auth.signIn returned %T", authorization)
}
signedUp, err := raw.AuthSignUp(ctx, &tg.AuthSignUpRequest{
PhoneNumber: record.Phone, PhoneCodeHash: sentCode.PhoneCodeHash, FirstName: record.FirstName,
})
if err != nil {
return fmt.Errorf("auth.signUp: %w", err)
}
authorized, ok := signedUp.(*tg.AuthAuthorization)
if !ok {
return fmt.Errorf("auth.signUp returned %T", signedUp)
}
user, ok = authorized.User.(*tg.User)
if !ok {
return fmt.Errorf("auth.signUp user is %T", authorized.User)
}
return nil
})
if err != nil {
return nil, err
}
if user == nil {
return nil, errors.New("provision completed without a user")
}
return user, nil
}
var _ session.Storage = (*EncryptedFileStorage)(nil)

View file

@ -0,0 +1,33 @@
package loadharness
import (
"path/filepath"
"testing"
)
func TestSessionDirectoryForManifestIsolatesNamedBundles(t *testing.T) {
tests := []struct {
manifest string
want string
}{
{manifest: filepath.Join("data", "load500", "manifest.json"), want: "sessions"},
{manifest: filepath.Join("data", "manifest-50.json"), want: "sessions-manifest-50"},
{manifest: filepath.Join("data", "manifest-500.json"), want: "sessions-manifest-500"},
}
for _, test := range tests {
t.Run(test.want, func(t *testing.T) {
if got := sessionDirectoryForManifest(test.manifest); got != test.want {
t.Fatalf("session directory = %q, want %q", got, test.want)
}
})
}
}
func TestDesiredSessionRecordUsesManifestNamespace(t *testing.T) {
cfg := ProvisionConfig{ManifestPath: filepath.Join("data", "manifest-500.json"), PhonePrefix: "+155500", FirstNamePrefix: "Load"}
record := desiredSessionRecord(12, 12, 1, cfg)
want := filepath.ToSlash(filepath.Join("sessions-manifest-500", "session-0012-device-1.bin"))
if record.SessionFile != want {
t.Fatalf("session file = %q, want %q", record.SessionFile, want)
}
}

View file

@ -0,0 +1,251 @@
package loadharness
import (
"encoding/json"
"fmt"
"math"
"os"
"path/filepath"
"sort"
"sync"
"sync/atomic"
"time"
)
var latencyBounds = [...]time.Duration{
5 * time.Millisecond, 10 * time.Millisecond, 25 * time.Millisecond,
50 * time.Millisecond, 100 * time.Millisecond, 250 * time.Millisecond,
500 * time.Millisecond, time.Second, 2 * time.Second, 5 * time.Second,
10 * time.Second, 30 * time.Second,
}
type operationMetrics struct {
count atomic.Uint64
errors atomic.Uint64
canceled atomic.Uint64
floodWaits atomic.Uint64
timeouts atomic.Uint64
connections atomic.Uint64
sumNS atomic.Int64
maxNS atomic.Int64
buckets [len(latencyBounds)]atomic.Uint64
}
func (m *operationMetrics) observe(start time.Time, err error) {
d := time.Since(start)
if d < 0 {
d = 0
}
m.count.Add(1)
m.sumNS.Add(int64(d))
for {
previous := m.maxNS.Load()
if int64(d) <= previous || m.maxNS.CompareAndSwap(previous, int64(d)) {
break
}
}
for i, bound := range latencyBounds {
if d <= bound {
m.buckets[i].Add(1)
}
}
if err != nil {
outcome := classifyError(err)
if outcome == "canceled" {
m.canceled.Add(1)
return
}
m.errors.Add(1)
switch outcome {
case "flood_wait":
m.floodWaits.Add(1)
case "timeout":
m.timeouts.Add(1)
case "connection":
m.connections.Add(1)
}
}
}
type OperationReport struct {
Count uint64 `json:"count"`
Errors uint64 `json:"errors"`
Canceled uint64 `json:"canceled"`
FloodWaits uint64 `json:"flood_waits"`
Timeouts uint64 `json:"timeouts"`
ConnectionErrors uint64 `json:"connection_errors"`
MeanMS float64 `json:"mean_ms"`
P50UpperMS float64 `json:"p50_upper_ms"`
P95UpperMS float64 `json:"p95_upper_ms"`
P99UpperMS float64 `json:"p99_upper_ms"`
MaxMS float64 `json:"max_ms"`
}
func (m *operationMetrics) report() OperationReport {
count := m.count.Load()
report := OperationReport{
Count: count, Errors: m.errors.Load(), Canceled: m.canceled.Load(), FloodWaits: m.floodWaits.Load(), Timeouts: m.timeouts.Load(), ConnectionErrors: m.connections.Load(),
MaxMS: durationMS(time.Duration(m.maxNS.Load())),
}
if count > 0 {
report.MeanMS = durationMS(time.Duration(m.sumNS.Load() / int64(count)))
report.P50UpperMS = durationMS(m.quantile(count, 0.50))
report.P95UpperMS = durationMS(m.quantile(count, 0.95))
report.P99UpperMS = durationMS(m.quantile(count, 0.99))
}
return report
}
func (m *operationMetrics) quantile(count uint64, q float64) time.Duration {
target := uint64(math.Ceil(float64(count) * q))
for i, bound := range latencyBounds {
if m.buckets[i].Load() >= target {
return bound
}
}
return latencyBounds[len(latencyBounds)-1]
}
func durationMS(d time.Duration) float64 {
return math.Round(float64(d)/float64(time.Millisecond)*1000) / 1000
}
type metricSet struct {
mu sync.RWMutex
ops map[string]*operationMetrics
}
func newMetricSet(names ...string) *metricSet {
m := &metricSet{ops: make(map[string]*operationMetrics, len(names))}
for _, name := range names {
m.ops[name] = &operationMetrics{}
}
return m
}
func (m *metricSet) observe(name string, start time.Time, err error) {
debugOperationError(name, err)
m.mu.RLock()
op := m.ops[name]
m.mu.RUnlock()
if op == nil {
// Operation names are code-owned and finite, but retain a lock-protected
// fallback for optional scenarios added by the harness.
m.mu.Lock()
op = m.ops[name]
if op == nil && len(m.ops) < 32 {
op = &operationMetrics{}
m.ops[name] = op
}
m.mu.Unlock()
}
if op != nil {
op.observe(start, err)
}
}
func (m *metricSet) report() map[string]OperationReport {
m.mu.RLock()
defer m.mu.RUnlock()
out := make(map[string]OperationReport, len(m.ops))
for name, op := range m.ops {
out[name] = op.report()
}
return out
}
type RunReport struct {
Version int `json:"version"`
StartedAt time.Time `json:"started_at"`
LoadEndedAt time.Time `json:"load_ended_at"`
FinishedAt time.Time `json:"finished_at"`
RequestedDuration string `json:"requested_duration"`
RecoveryDuration string `json:"recovery_duration"`
ExpectedSessions int `json:"expected_sessions"`
PeakReadySessions int `json:"peak_ready_sessions"`
FinalReadySessions int `json:"final_ready_sessions"`
SteadySamples int `json:"steady_samples"`
SteadyReadyRatio float64 `json:"steady_ready_ratio"`
MinSteadyReadySessions int `json:"min_steady_ready_sessions"`
ConnectionAttempts uint64 `json:"connection_attempts"`
Reconnects uint64 `json:"reconnects"`
Disconnects uint64 `json:"disconnects"`
UpdatesReceived uint64 `json:"updates_received"`
DownloadedBytes uint64 `json:"downloaded_bytes"`
WorkerFatalErrors uint64 `json:"worker_fatal_errors"`
Operations map[string]OperationReport `json:"operations"`
BaselineServerMetrics map[string]float64 `json:"baseline_server_metrics,omitempty"`
FinalServerMetrics map[string]float64 `json:"final_server_metrics,omitempty"`
ServerMetricsScrapes uint64 `json:"server_metrics_scrapes"`
ServerMetricsErrors uint64 `json:"server_metrics_errors"`
Pass bool `json:"pass"`
Failures []string `json:"failures,omitempty"`
}
func WriteReport(path string, report *RunReport) error {
data, err := json.MarshalIndent(report, "", " ")
if err != nil {
return fmt.Errorf("encode report: %w", err)
}
return writeFileAtomic(path, append(data, '\n'), 0o600)
}
type eventWriter struct {
mu sync.Mutex
f *os.File
written uint64
dropped uint64
}
func newEventWriter(path string) (*eventWriter, error) {
if path == "" {
return &eventWriter{}, nil
}
if err := os.MkdirAll(filepath.Dir(path), 0o700); err != nil {
return nil, err
}
f, err := os.OpenFile(path, os.O_CREATE|os.O_WRONLY|os.O_TRUNC, 0o600)
if err != nil {
return nil, err
}
return &eventWriter{f: f}, nil
}
func (w *eventWriter) write(value any) {
if w == nil || w.f == nil {
return
}
data, err := json.Marshal(value)
if err != nil {
return
}
w.mu.Lock()
if w.written >= 10000 {
w.dropped++
w.mu.Unlock()
return
}
_, _ = w.f.Write(append(data, '\n'))
w.written++
w.mu.Unlock()
}
func (w *eventWriter) close() error {
if w == nil || w.f == nil {
return nil
}
w.mu.Lock()
err := w.f.Close()
w.f = nil
w.mu.Unlock()
return err
}
func sortedOperationNames(ops map[string]OperationReport) []string {
names := make([]string, 0, len(ops))
for name := range ops {
names = append(names, name)
}
sort.Strings(names)
return names
}

View file

@ -0,0 +1,127 @@
package loadharness
import (
"context"
"errors"
"fmt"
"net"
"testing"
"time"
"github.com/iamxvbaba/td/pool"
tdrpc "github.com/iamxvbaba/td/rpc"
)
func TestOperationMetricsUsesBoundedHistogramAndFixedErrorClasses(t *testing.T) {
metrics := &operationMetrics{}
metrics.observe(time.Now().Add(-20*time.Millisecond), nil)
metrics.observe(time.Now().Add(-200*time.Millisecond), errors.New("FLOOD_WAIT_1 phone=secret"))
report := metrics.report()
if report.Count != 2 || report.Errors != 1 || report.FloodWaits != 1 {
t.Fatalf("report = %#v", report)
}
if report.P50UpperMS <= 0 || report.P99UpperMS < report.P50UpperMS || report.MaxMS <= 0 {
t.Fatalf("latency report = %#v", report)
}
}
func TestClassifyErrorReasonUsesFiniteRedactedVocabulary(t *testing.T) {
tests := []struct {
err error
want string
}{
{errors.New("dial tcp 10.0.0.1:2398: socket: too many open files"), "file_descriptor_limit"},
{errors.New("read: temporary auth key not found: pfs reconnect required"), "pfs_reconnect"},
{errors.New("read tcp: EOF auth_key_id=secret"), "eof"},
}
for _, test := range tests {
if got := classifyErrorReason(test.err); got != test.want {
t.Fatalf("classifyErrorReason(%v) = %q, want %q", test.err, got, test.want)
}
}
}
func TestClassifyErrorRecognizesTypedReconnectFailures(t *testing.T) {
tests := []error{
fmt.Errorf("invoke: %w", tdrpc.ErrEngineClosed),
fmt.Errorf("acquire: %w", pool.ErrConnDead),
fmt.Errorf("read: %w", net.ErrClosed),
errors.New("write: broken pipe"),
}
for _, err := range tests {
if got := classifyError(err); got != "connection" {
t.Fatalf("classifyError(%v) = %q, want connection", err, got)
}
}
}
func TestOperationMetricsSeparatesHarnessCancellation(t *testing.T) {
metrics := &operationMetrics{}
metrics.observe(time.Now(), context.Canceled)
report := metrics.report()
if report.Count != 1 || report.Canceled != 1 || report.Errors != 0 {
t.Fatalf("report = %#v", report)
}
}
func TestEvaluateReportAllowsOnlyConnectionErrorsForExpectedRestart(t *testing.T) {
report := &RunReport{
ExpectedSessions: 2, PeakReadySessions: 2, Reconnects: 2,
SteadySamples: 1, SteadyReadyRatio: 1, MinSteadyReadySessions: 2,
Operations: map[string]OperationReport{
"connection.dead": {Count: 2, Errors: 2, ConnectionErrors: 2},
},
}
evaluateReport(report, RunConfig{MinimumReadyRatio: 1, ExpectServerRestart: true})
if !report.Pass {
t.Fatalf("report = %#v", report)
}
report.Operations["ping"] = OperationReport{Count: 1, Errors: 1}
report.Failures = nil
evaluateReport(report, RunConfig{MinimumReadyRatio: 1, ExpectServerRestart: true})
if report.Pass {
t.Fatalf("unexpected application error passed: %#v", report)
}
}
func TestEvaluateReportRequiresReclamationAndNoFloodWait(t *testing.T) {
report := &RunReport{
ExpectedSessions: 10, PeakReadySessions: 10, ServerMetricsScrapes: 1,
SteadySamples: 1, SteadyReadyRatio: 1, MinSteadyReadySessions: 10,
Operations: map[string]OperationReport{"ping": {Count: 10}},
BaselineServerMetrics: map[string]float64{
"telesrv_mtproto_raw_connections": 2,
"telesrv_mtproto_logical_outbox_bytes": 3,
},
FinalServerMetrics: map[string]float64{
"telesrv_mtproto_raw_connections": 2,
"telesrv_mtproto_logical_outbox_bytes": 4,
},
}
evaluateReport(report, RunConfig{MinimumReadyRatio: 1, RecoveryDuration: time.Minute, ServerMetricsURL: "http://metrics"})
if report.Pass || len(report.Failures) != 1 {
t.Fatalf("report = %#v", report)
}
}
func TestEvaluateReportAcceptsReturnToNonZeroSharedServerBaseline(t *testing.T) {
report := &RunReport{
ExpectedSessions: 10, PeakReadySessions: 10, ServerMetricsScrapes: 2,
SteadySamples: 1, SteadyReadyRatio: 1, MinSteadyReadySessions: 10,
Operations: map[string]OperationReport{"ping": {Count: 10}},
BaselineServerMetrics: map[string]float64{
"telesrv_mtproto_raw_connections": 2,
"telesrv_mtproto_logical_sessions": 2,
"telesrv_mtproto_logical_outbox_bytes": 1024,
},
FinalServerMetrics: map[string]float64{
"telesrv_mtproto_raw_connections": 2,
"telesrv_mtproto_logical_sessions": 2,
"telesrv_mtproto_logical_outbox_bytes": 1024,
},
}
evaluateReport(report, RunConfig{MinimumReadyRatio: 1, RecoveryDuration: time.Minute, ServerMetricsURL: "http://metrics"})
if !report.Pass || len(report.Failures) != 0 {
t.Fatalf("report = %#v", report)
}
}

1044
internal/loadharness/run.go Normal file

File diff suppressed because it is too large Load diff

View file

@ -0,0 +1,140 @@
package loadharness
import (
"bufio"
"context"
"fmt"
"io"
"math"
"net/http"
"strconv"
"strings"
"sync/atomic"
"time"
)
const maxServerMetricsBytes = 4 << 20
var selectedServerMetrics = map[string]struct{}{
"telesrv_mtproto_raw_connections": {},
"telesrv_mtproto_sessions": {},
"telesrv_mtproto_logical_sessions": {},
"telesrv_mtproto_logical_outbox_frames": {},
"telesrv_mtproto_logical_outbox_bytes": {},
"telesrv_mtproto_logical_outbox_acked_frames_total": {},
"telesrv_mtproto_logical_outbox_acked_bytes_total": {},
"telesrv_mtproto_logical_outbox_retained_seconds_count": {},
"telesrv_mtproto_logical_outbox_retained_seconds_sum": {},
"telesrv_mtproto_pending_push_bytes": {},
"telesrv_mtproto_inbound_rpc_tasks": {},
"telesrv_mtproto_inbound_rpc_bytes": {},
"telesrv_mtproto_inbound_frame_bytes": {},
"telesrv_mtproto_outbound_tracked_bytes": {},
"telesrv_mtproto_outbound_write_bytes": {},
"telesrv_mtproto_rpc_result_owners": {},
"telesrv_mtproto_rpc_result_receipts": {},
"telesrv_mtproto_rpc_result_receipt_bytes": {},
"telesrv_mtproto_rpc_result_subscribers": {},
"telesrv_mtproto_rpc_result_inner_bytes_total": {},
"telesrv_mtproto_rpc_result_wire_bytes_total": {},
"telesrv_mtproto_rpc_result_delivered_bytes_total": {},
"telesrv_go_goroutines": {},
"telesrv_go_heap_alloc_bytes": {},
"telesrv_go_heap_inuse_bytes": {},
"telesrv_go_heap_objects": {},
"telesrv_go_sys_bytes": {},
"telesrv_postgres_pool_connections": {},
"telesrv_postgres_pool_acquire_wait_seconds": {},
"telesrv_postgres_pool_empty_acquire_count": {},
"telesrv_postgres_pool_canceled_acquire_count": {},
"telesrv_redis_pool_connections": {},
"telesrv_redis_pool_pending_requests": {},
"telesrv_redis_pool_timeouts": {},
"telesrv_redis_pool_wait_seconds": {},
"telesrv_metrics_dropped_observations_total": {},
}
type serverMetricsClient struct {
url string
client *http.Client
success atomic.Uint64
errors atomic.Uint64
}
func newServerMetricsClient(url string) *serverMetricsClient {
if strings.TrimSpace(url) == "" {
return nil
}
return &serverMetricsClient{url: url, client: &http.Client{Timeout: 5 * time.Second}}
}
func (c *serverMetricsClient) scrape(ctx context.Context) (map[string]float64, error) {
if c == nil {
return nil, nil
}
request, err := http.NewRequestWithContext(ctx, http.MethodGet, c.url, nil)
if err != nil {
c.errors.Add(1)
return nil, err
}
response, err := c.client.Do(request)
if err != nil {
c.errors.Add(1)
return nil, err
}
defer response.Body.Close()
if response.StatusCode != http.StatusOK {
c.errors.Add(1)
return nil, fmt.Errorf("metrics HTTP status %d", response.StatusCode)
}
reader := bufio.NewScanner(io.LimitReader(response.Body, maxServerMetricsBytes))
reader.Buffer(make([]byte, 64<<10), 1<<20)
values := make(map[string]float64, len(selectedServerMetrics))
for reader.Scan() {
line := strings.TrimSpace(reader.Text())
if line == "" || strings.HasPrefix(line, "#") {
continue
}
fields := strings.Fields(line)
if len(fields) < 2 {
continue
}
name := fields[0]
if idx := strings.IndexByte(name, '{'); idx >= 0 {
name = name[:idx]
}
if _, ok := selectedServerMetrics[name]; !ok {
continue
}
value, err := strconv.ParseFloat(fields[1], 64)
if err != nil || math.IsNaN(value) || math.IsInf(value, 0) {
continue
}
// Reports need bounded, comparable capacity signals, not an unbounded copy
// of Prometheus label series. Aggregate every selected family into one
// key so method/encoding cardinality can never starve later gauges (the
// endpoint orders counters before gauges). The source /metrics endpoint
// retains full labels for detailed diagnosis.
values[name] += value
}
if err := reader.Err(); err != nil {
c.errors.Add(1)
return nil, err
}
c.success.Add(1)
return values, nil
}
func (c *serverMetricsClient) successes() uint64 {
if c == nil {
return 0
}
return c.success.Load()
}
func (c *serverMetricsClient) failures() uint64 {
if c == nil {
return 0
}
return c.errors.Load()
}

View file

@ -0,0 +1,33 @@
package loadharness
import (
"context"
"fmt"
"net/http"
"net/http/httptest"
"testing"
)
func TestServerMetricsScrapeSelectsBoundedCapacitySignals(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
fmt.Fprintln(w, `telesrv_mtproto_raw_connections 500`)
fmt.Fprintln(w, `telesrv_mtproto_sessions{state="active"} 499`)
fmt.Fprintln(w, `telesrv_mtproto_sessions{state="provisional"} 1`)
for i := 0; i < 256; i++ {
fmt.Fprintf(w, "telesrv_mtproto_rpc_result_wire_bytes_total{method=%q} 1\n", fmt.Sprintf("method-%d", i))
}
fmt.Fprintln(w, `unrelated_high_cardinality{user_id="secret"} 1`)
}))
defer server.Close()
client := newServerMetricsClient(server.URL)
values, err := client.scrape(context.Background())
if err != nil {
t.Fatal(err)
}
if values["telesrv_mtproto_raw_connections"] != 500 || values["telesrv_mtproto_sessions"] != 500 || values["telesrv_mtproto_rpc_result_wire_bytes_total"] != 256 {
t.Fatalf("values = %#v", values)
}
if len(values) != 3 || client.successes() != 1 || client.failures() != 0 {
t.Fatalf("bounded values/scrapes = %#v, %d/%d", values, client.successes(), client.failures())
}
}

View file

@ -0,0 +1,165 @@
package loadharness
import (
"context"
"crypto/aes"
"crypto/cipher"
"crypto/rand"
"encoding/base64"
"errors"
"fmt"
"io"
"os"
"path/filepath"
"runtime"
"strings"
"sync"
"github.com/iamxvbaba/td/session"
)
const encryptedSessionMagic = "TLSLOAD1"
// EncryptedFileStorage encrypts gotd's complete session blob with AES-256-GCM.
// A unique random nonce is generated on every replacement and the file is
// written with owner-only permissions.
type EncryptedFileStorage struct {
Path string
Key [32]byte
mu sync.Mutex
}
func (s *EncryptedFileStorage) LoadSession(context.Context) ([]byte, error) {
if s == nil || strings.TrimSpace(s.Path) == "" {
return nil, errors.New("invalid encrypted session storage")
}
s.mu.Lock()
defer s.mu.Unlock()
data, err := os.ReadFile(s.Path)
if os.IsNotExist(err) {
return nil, session.ErrNotFound
}
if err != nil {
return nil, fmt.Errorf("read encrypted session: %w", err)
}
block, err := aes.NewCipher(s.Key[:])
if err != nil {
return nil, err
}
gcm, err := cipher.NewGCM(block)
if err != nil {
return nil, err
}
header := len(encryptedSessionMagic) + gcm.NonceSize()
if len(data) < header || string(data[:len(encryptedSessionMagic)]) != encryptedSessionMagic {
return nil, errors.New("encrypted session has an invalid header")
}
nonce := data[len(encryptedSessionMagic):header]
plain, err := gcm.Open(nil, nonce, data[header:], []byte(encryptedSessionMagic))
if err != nil {
return nil, errors.New("encrypted session authentication failed")
}
return plain, nil
}
func (s *EncryptedFileStorage) StoreSession(_ context.Context, plain []byte) error {
if s == nil || strings.TrimSpace(s.Path) == "" {
return errors.New("invalid encrypted session storage")
}
s.mu.Lock()
defer s.mu.Unlock()
block, err := aes.NewCipher(s.Key[:])
if err != nil {
return err
}
gcm, err := cipher.NewGCM(block)
if err != nil {
return err
}
nonce := make([]byte, gcm.NonceSize())
if _, err := io.ReadFull(rand.Reader, nonce); err != nil {
return fmt.Errorf("generate session nonce: %w", err)
}
data := make([]byte, 0, len(encryptedSessionMagic)+len(nonce)+len(plain)+gcm.Overhead())
data = append(data, encryptedSessionMagic...)
data = append(data, nonce...)
data = gcm.Seal(data, nonce, plain, []byte(encryptedSessionMagic))
return writeFileAtomic(s.Path, data, 0o600)
}
func GenerateSessionKey(path string) error {
if _, err := os.Stat(path); err == nil {
return fmt.Errorf("refusing to overwrite existing session key %q", path)
} else if !os.IsNotExist(err) {
return err
}
var key [32]byte
if _, err := io.ReadFull(rand.Reader, key[:]); err != nil {
return err
}
encoded := base64.StdEncoding.EncodeToString(key[:]) + "\n"
return writeFileAtomic(path, []byte(encoded), 0o600)
}
func LoadSessionKey(path string) ([32]byte, error) {
var key [32]byte
info, err := os.Stat(path)
if err != nil {
return key, fmt.Errorf("stat session key: %w", err)
}
if runtime.GOOS != "windows" && info.Mode().Perm()&0o077 != 0 {
return key, fmt.Errorf("session key %q must not be group/world accessible (mode %o)", path, info.Mode().Perm())
}
data, err := os.ReadFile(path)
if err != nil {
return key, err
}
decoded, err := base64.StdEncoding.DecodeString(strings.TrimSpace(string(data)))
if err != nil || len(decoded) != len(key) {
return key, errors.New("session key must be base64-encoded 32 bytes")
}
copy(key[:], decoded)
return key, nil
}
func writeFileAtomic(path string, data []byte, mode os.FileMode) (retErr error) {
dir := filepath.Dir(path)
if err := os.MkdirAll(dir, 0o700); err != nil {
return err
}
tmp, err := os.CreateTemp(dir, ".telesrv-load-*")
if err != nil {
return err
}
tmpName := tmp.Name()
defer func() {
_ = tmp.Close()
if retErr != nil {
_ = os.Remove(tmpName)
}
}()
if err := tmp.Chmod(mode); err != nil {
return err
}
if _, err := tmp.Write(data); err != nil {
return err
}
if err := tmp.Sync(); err != nil {
return err
}
if err := tmp.Close(); err != nil {
return err
}
// On Unix rename atomically replaces. Windows requires removing the old
// destination first; session files remain recoverable from the complete temp
// file if that narrow replacement fails.
if runtime.GOOS == "windows" {
if err := os.Remove(path); err != nil && !os.IsNotExist(err) {
return err
}
}
if err := os.Rename(tmpName, path); err != nil {
return err
}
return nil
}

View file

@ -0,0 +1,88 @@
package loadharness
import (
"bytes"
"context"
"os"
"path/filepath"
"runtime"
"testing"
)
func TestEncryptedFileStorageRoundTripAndNonceRotation(t *testing.T) {
dir := t.TempDir()
path := filepath.Join(dir, "session.bin")
var key [32]byte
for i := range key {
key[i] = byte(i + 1)
}
storage := &EncryptedFileStorage{Path: path, Key: key}
plain := []byte(`{"auth_key":"plaintext-secret-marker"}`)
if err := storage.StoreSession(context.Background(), plain); err != nil {
t.Fatal(err)
}
first, err := os.ReadFile(path)
if err != nil {
t.Fatal(err)
}
if bytes.Contains(first, []byte("plaintext-secret-marker")) {
t.Fatal("encrypted session retained plaintext auth material")
}
if got, err := storage.LoadSession(context.Background()); err != nil || !bytes.Equal(got, plain) {
t.Fatalf("round trip = %q, %v", got, err)
}
if err := storage.StoreSession(context.Background(), plain); err != nil {
t.Fatal(err)
}
second, err := os.ReadFile(path)
if err != nil {
t.Fatal(err)
}
if bytes.Equal(first, second) {
t.Fatal("successive session writes reused ciphertext/nonce")
}
wrong := key
wrong[0] ^= 0xff
if _, err := (&EncryptedFileStorage{Path: path, Key: wrong}).LoadSession(context.Background()); err == nil {
t.Fatal("wrong session key unexpectedly authenticated")
}
if runtime.GOOS != "windows" {
info, err := os.Stat(path)
if err != nil {
t.Fatal(err)
}
if got := info.Mode().Perm(); got != 0o600 {
t.Fatalf("session mode = %o, want 600", got)
}
}
}
func TestSessionKeyGenerationRefusesOverwrite(t *testing.T) {
path := filepath.Join(t.TempDir(), "session.key")
if err := GenerateSessionKey(path); err != nil {
t.Fatal(err)
}
first, err := LoadSessionKey(path)
if err != nil {
t.Fatal(err)
}
if first == ([32]byte{}) {
t.Fatal("generated all-zero key")
}
if err := GenerateSessionKey(path); err == nil {
t.Fatal("keygen overwrote an existing key")
}
}
func TestWriteFileAtomicReplacesExisting(t *testing.T) {
path := filepath.Join(t.TempDir(), "report.json")
if err := writeFileAtomic(path, []byte("first"), 0o600); err != nil {
t.Fatal(err)
}
if err := writeFileAtomic(path, []byte("second"), 0o600); err != nil {
t.Fatal(err)
}
if got, err := os.ReadFile(path); err != nil || string(got) != "second" {
t.Fatalf("replacement = %q, %v", got, err)
}
}

View file

@ -0,0 +1,140 @@
// Package loadharness implements the real-MTProto capacity harness used by
// cmd/telesrv-load. It deliberately uses the published gotd fork instead of
// server-internal handlers or direct database fixtures.
package loadharness
import (
"encoding/json"
"errors"
"fmt"
"os"
"path/filepath"
"sort"
"strings"
"time"
)
const ManifestVersion = 1
// Endpoint is the immutable wire target shared by provisioning and runs.
type Endpoint struct {
Address string `json:"address"`
DC int `json:"dc"`
APIID int `json:"api_id"`
APIHash string `json:"api_hash"`
RSAKeyPath string `json:"rsa_key_path"`
Obfuscated bool `json:"obfuscated"`
PFS bool `json:"pfs"`
TempKeyTTL int `json:"temp_key_ttl_seconds,omitempty"`
}
// SessionRecord maps one physical MTProto session file to one logical account.
// It contains routing facts only; auth key material remains in encrypted files.
type SessionRecord struct {
Index int `json:"index"`
AccountIndex int `json:"account_index"`
DeviceIndex int `json:"device_index"`
Phone string `json:"phone"`
FirstName string `json:"first_name"`
SessionFile string `json:"session_file"`
UserID int64 `json:"user_id"`
AccessHash int64 `json:"access_hash"`
}
// Manifest never embeds session encryption keys, auth keys, phone-code hashes
// or raw server errors. It does contain generated test phone/user routing data,
// so it remains a controlled run artifact and is not copied into RunReport.
type Manifest struct {
Version int `json:"version"`
CreatedAt time.Time `json:"created_at"`
Endpoint Endpoint `json:"endpoint"`
Sessions []SessionRecord `json:"sessions"`
}
func (e Endpoint) Validate() error {
if strings.TrimSpace(e.Address) == "" {
return errors.New("endpoint address is required")
}
if e.DC == 0 {
return errors.New("endpoint DC must be non-zero")
}
if e.APIID <= 0 || strings.TrimSpace(e.APIHash) == "" {
return errors.New("endpoint api_id and api_hash are required")
}
if strings.TrimSpace(e.RSAKeyPath) == "" {
return errors.New("endpoint RSA key path is required")
}
return nil
}
func (m *Manifest) Validate() error {
if m == nil {
return errors.New("nil manifest")
}
if m.Version != ManifestVersion {
return fmt.Errorf("manifest version %d, want %d", m.Version, ManifestVersion)
}
if err := m.Endpoint.Validate(); err != nil {
return err
}
indices := make(map[int]struct{}, len(m.Sessions))
files := make(map[string]struct{}, len(m.Sessions))
for _, session := range m.Sessions {
if session.Index < 0 || session.AccountIndex < 0 || session.DeviceIndex < 0 {
return fmt.Errorf("session %d has a negative index", session.Index)
}
if _, ok := indices[session.Index]; ok {
return fmt.Errorf("duplicate session index %d", session.Index)
}
indices[session.Index] = struct{}{}
if strings.TrimSpace(session.Phone) == "" || strings.TrimSpace(session.SessionFile) == "" {
return fmt.Errorf("session %d is missing phone or session_file", session.Index)
}
clean := filepath.Clean(session.SessionFile)
if filepath.IsAbs(clean) || clean == "." || strings.HasPrefix(clean, ".."+string(filepath.Separator)) || clean == ".." {
return fmt.Errorf("session %d has unsafe session_file %q", session.Index, session.SessionFile)
}
if _, ok := files[clean]; ok {
return fmt.Errorf("duplicate session file %q", clean)
}
files[clean] = struct{}{}
if session.UserID <= 0 {
return fmt.Errorf("session %d has no provisioned user_id", session.Index)
}
}
return nil
}
func LoadManifest(path string) (*Manifest, error) {
data, err := os.ReadFile(path)
if err != nil {
return nil, fmt.Errorf("read manifest: %w", err)
}
var manifest Manifest
decoder := json.NewDecoder(strings.NewReader(string(data)))
decoder.DisallowUnknownFields()
if err := decoder.Decode(&manifest); err != nil {
return nil, fmt.Errorf("decode manifest: %w", err)
}
if err := manifest.Validate(); err != nil {
return nil, err
}
sort.Slice(manifest.Sessions, func(i, j int) bool { return manifest.Sessions[i].Index < manifest.Sessions[j].Index })
return &manifest, nil
}
func WriteManifest(path string, manifest *Manifest) error {
if err := manifest.Validate(); err != nil {
return err
}
data, err := json.MarshalIndent(manifest, "", " ")
if err != nil {
return fmt.Errorf("encode manifest: %w", err)
}
data = append(data, '\n')
return writeFileAtomic(path, data, 0o600)
}
func resolveSessionPath(manifestPath string, record SessionRecord) string {
return filepath.Join(filepath.Dir(manifestPath), filepath.FromSlash(record.SessionFile))
}

View file

@ -0,0 +1,74 @@
package loadharness
import (
"path/filepath"
"testing"
"time"
)
func validManifest() *Manifest {
return &Manifest{
Version: ManifestVersion, CreatedAt: time.Now(),
Endpoint: Endpoint{Address: "127.0.0.1:2398", DC: 2, APIID: 1, APIHash: "hash", RSAKeyPath: "server.pem"},
Sessions: []SessionRecord{{
Index: 0, AccountIndex: 0, DeviceIndex: 0, Phone: "+155500000001", FirstName: "Load0001",
SessionFile: "sessions/session-0000-device-0.bin", UserID: 1, AccessHash: 2,
}},
}
}
func TestManifestRoundTripContainsNoSessionSecrets(t *testing.T) {
path := filepath.Join(t.TempDir(), "manifest.json")
manifest := validManifest()
if err := WriteManifest(path, manifest); err != nil {
t.Fatal(err)
}
loaded, err := LoadManifest(path)
if err != nil {
t.Fatal(err)
}
if len(loaded.Sessions) != 1 || loaded.Sessions[0].UserID != 1 {
t.Fatalf("loaded manifest = %#v", loaded)
}
}
func TestManifestRejectsEscapingAndDuplicateSessionPaths(t *testing.T) {
manifest := validManifest()
manifest.Sessions[0].SessionFile = "../outside.bin"
if err := manifest.Validate(); err == nil {
t.Fatal("escaping session path accepted")
}
manifest = validManifest()
duplicate := manifest.Sessions[0]
duplicate.Index = 1
duplicate.AccountIndex = 1
duplicate.UserID = 2
manifest.Sessions = append(manifest.Sessions, duplicate)
if err := manifest.Validate(); err == nil {
t.Fatal("duplicate session path accepted")
}
}
func TestExplicitZeroExtraDevicesAndRecoveryAreValid(t *testing.T) {
provision := ProvisionConfig{
ManifestPath: "manifest.json", SessionKeyPath: "key", RSAKeyPath: "rsa",
Endpoint: *&validManifest().Endpoint, Accounts: 1, ExtraDevices: 0, Concurrency: 1,
PhonePrefix: "+155500", Code: "12345", FirstNamePrefix: "Load",
}
if err := provision.validate(); err != nil {
t.Fatalf("zero extra devices: %v", err)
}
run := RunConfig{
ManifestPath: "manifest.json", SessionKeyPath: "key", ReportPath: "report.json",
Duration: time.Second, RecoveryDuration: 0, RampDuration: 0,
RPCInterval: time.Millisecond, MessageInterval: -1, SampleInterval: time.Millisecond,
OperationTimeout: time.Second, MinimumReadyRatio: 1,
}
if err := run.validate(); err != nil {
t.Fatalf("zero recovery/ramp: %v", err)
}
run.OperationTimeout = 0
if err := run.validate(); err == nil {
t.Fatal("zero operation timeout accepted")
}
}

View file

@ -76,9 +76,13 @@ func TestTelegramClientEndToEnd(t *testing.T) {
if cfg.ThisDC != dc {
t.Errorf("config.ThisDC = %d, want %d", cfg.ThisDC, dc)
}
// 不下发 DCOptions客户端使用自己的 DCList / 写死 static 地址。
if len(cfg.DCOptions) != 0 {
t.Errorf("config.DCOptions = %+v, want empty", cfg.DCOptions)
if len(cfg.DCOptions) != 1 {
t.Errorf("config.DCOptions = %+v, want one reconnect route", cfg.DCOptions)
} else {
option := cfg.DCOptions[0]
if option.ID != dc || option.IPAddress != tcpAddr.IP.String() || option.Port != tcpAddr.Port {
t.Errorf("config.DCOptions[0] = %+v, want dc=%d at %s", option, dc, tcpAddr)
}
}
return nil
}); err != nil {

View file

@ -2,8 +2,8 @@ package mtprotoedge
import "time"
// Metrics 接收连接层运行指标。实现可对接 Prometheus 等监控系统
// 默认 NopMetrics零开销。第一阶段仅预留钩子正式指标后续接入
// Metrics 接收连接层运行指标。生产入口接入有界 Prometheus exporter
// 其它 embedder 可继续使用 NopMetrics零开销
type Metrics interface {
// ConnOpened 在接受一个连接时调用。
ConnOpened()
@ -39,6 +39,14 @@ type RPCResultMetrics interface {
RPCResultDelivered(method string, egressLatency time.Duration, wireBytes int, err error)
}
// LogicalOutboxMetrics observes the sole owner of unacknowledged server frames.
// It is intentionally optional: embedders can keep the small Metrics surface,
// while production capacity tests can distinguish physical delivery from the
// later client ACK that actually releases retained bytes.
type LogicalOutboxMetrics interface {
LogicalOutboxAcknowledged(bytes int, retainedFor time.Duration, rpcResult bool)
}
// ConnectionIntakeMetrics is an optional extension for the pre-session
// connection pipeline. stage is one of raw_accept, mux_sniff, mux_delivery,
// transport_dispatch, transport_promote, or first_frame; outcome is a bounded

View file

@ -1834,13 +1834,13 @@ func (c *Conn) handleOutboundOp(state *outboundState, op outboundOp) {
state.mu.Lock()
var (
result outboundResult
acked []int64
acked []outboundAcknowledgement
)
switch op.kind {
case outboundSend:
result.err = c.handleOutboundSend(state, op)
case outboundAck:
acked = state.ack(op.ids)
acked = state.ackWithDetails(op.ids)
case outboundQueryState:
result.info = state.stateInfo(op.ids)
case outboundResend:
@ -1851,9 +1851,16 @@ func (c *Conn) handleOutboundOp(state *outboundState, op outboundOp) {
result.err = fmt.Errorf("unknown outbound op %d", op.kind)
}
state.mu.Unlock()
for _, reqMsgID := range acked {
if c.rpcResultAcked != nil {
c.rpcResultAcked(c, reqMsgID)
for _, ack := range acked {
if metrics, ok := c.metrics.(LogicalOutboxMetrics); ok {
retainedFor := time.Duration(0)
if !ack.sentAt.IsZero() {
retainedFor = time.Since(ack.sentAt)
}
metrics.LogicalOutboxAcknowledged(ack.bytes, retainedFor, ack.reqMsgID != 0)
}
if ack.reqMsgID != 0 && c.rpcResultAcked != nil {
c.rpcResultAcked(c, ack.reqMsgID)
}
}
op.finish(result)
@ -2663,25 +2670,45 @@ func (s *outboundState) addReserved(frame *outboundFrame) int {
return s.shrinkPending()
}
type outboundAcknowledgement struct {
reqMsgID int64
bytes int
sentAt time.Time
}
func (s *outboundState) ack(ids []int64) []int64 {
var requestIDs []int64
details := s.ackWithDetails(ids)
requestIDs := make([]int64, 0, len(details))
for _, detail := range details {
if detail.reqMsgID != 0 {
requestIDs = append(requestIDs, detail.reqMsgID)
}
}
return requestIDs
}
func (s *outboundState) ackWithDetails(ids []int64) []outboundAcknowledgement {
var acknowledged []outboundAcknowledgement
for _, id := range ids {
frame, ok := s.pending[id]
if !ok {
continue
}
if frame.reqMsgID != 0 {
requestIDs = append(requestIDs, frame.reqMsgID)
detail := outboundAcknowledgement{
reqMsgID: frame.reqMsgID,
bytes: len(frame.body),
sentAt: frame.sentAt,
}
if !s.removePending(id) {
continue
}
s.markAcked(id)
acknowledged = append(acknowledged, detail)
}
if len(s.order) > s.maxMessages*2 {
s.compactOrder()
}
return requestIDs
return acknowledged
}
func (s *outboundState) stateInfo(ids []int64) []byte {

View file

@ -29,6 +29,21 @@ type failAfterTransport struct {
last []byte
}
type acknowledgementCaptureMetrics struct {
NopMetrics
count atomic.Int64
bytes atomic.Int64
retainedNS atomic.Int64
rpcResult atomic.Bool
}
func (m *acknowledgementCaptureMetrics) LogicalOutboxAcknowledged(bytes int, retainedFor time.Duration, rpcResult bool) {
m.count.Add(1)
m.bytes.Add(int64(bytes))
m.retainedNS.Store(int64(retainedFor))
m.rpcResult.Store(rpcResult)
}
func TestRPCResultReplayAttemptHooksArePhysicalConnectionLocal(t *testing.T) {
const reqMsgID = int64(771)
base := &encodedOutboundMessage{
@ -814,6 +829,8 @@ func TestOutboundTrackedBudgetAckAndCloseReturnExactly(t *testing.T) {
budget := newOutboundTrackedBudget(64)
tr := &failAfterTransport{}
c := newOutboundTestConn(t, tr, budget)
metrics := &acknowledgementCaptureMetrics{}
c.metrics = metrics
ctx, cancel := context.WithTimeout(context.Background(), time.Second)
defer cancel()
body := exactTestUpdatesEncoded(t, c, make([]byte, 12))
@ -827,6 +844,9 @@ func TestOutboundTrackedBudgetAckAndCloseReturnExactly(t *testing.T) {
if err != nil {
t.Fatalf("decrypt frame: %v", err)
}
// Windows wall-clock resolution can otherwise make an immediate ACK look
// like zero retention even though sentAt was populated after the write.
time.Sleep(time.Millisecond)
c.AckServerMessages([]int64{data.MessageID})
deadline := time.Now().Add(time.Second)
for budget.snapshot() != 0 && time.Now().Before(deadline) {
@ -835,6 +855,18 @@ func TestOutboundTrackedBudgetAckAndCloseReturnExactly(t *testing.T) {
if got := budget.snapshot(); got != 0 {
t.Fatalf("tracked bytes after ack = %d, want 0", got)
}
if got := metrics.count.Load(); got != 1 {
t.Fatalf("logical ACK metric count = %d, want 1", got)
}
if got := metrics.bytes.Load(); got != 12 {
t.Fatalf("logical ACK metric bytes = %d, want 12", got)
}
if metrics.retainedNS.Load() <= 0 {
t.Fatal("logical ACK metric did not record positive retention")
}
if metrics.rpcResult.Load() {
t.Fatal("ordinary update ACK was classified as rpc_result")
}
})
t.Run("close", func(t *testing.T) {

View file

@ -62,9 +62,12 @@ func TestRPCGetConfig(t *testing.T) {
if cfg.ThisDC != dc {
t.Fatalf("config.ThisDC = %d, want %d", cfg.ThisDC, dc)
}
// 不下发 DCOptions客户端使用写死的 static DC 地址(空列表令其保留本地地址)。
if len(cfg.DCOptions) != 0 {
t.Fatalf("config.DCOptions = %+v, want empty (client uses pinned static address)", cfg.DCOptions)
if len(cfg.DCOptions) != 1 {
t.Fatalf("config.DCOptions = %+v, want one reconnect route", cfg.DCOptions)
}
option := cfg.DCOptions[0]
if option.ID != dc || option.IPAddress != advIP || option.Port != advPort {
t.Fatalf("config.DCOptions[0] = %+v, want dc=%d at %s:%d", option, dc, advIP, advPort)
}
}

View file

@ -0,0 +1,186 @@
package mtprotoedge
// RuntimeSnapshot is a point-in-time, identity-free view of the MTProto edge.
// It deliberately exposes only bounded aggregate values so callers can publish
// it through a metrics endpoint without leaking auth keys, sessions or remote
// addresses. Values from independently locked components can differ by one
// concurrent transition; every individual budget/count remains internally
// consistent.
type RuntimeSnapshot struct {
RawConnections int64
RawConnectionLimit int64
Handshakes int64
HandshakeLimit int64
ActiveSessions int64
ProvisionalSessions int64
LogicalSessions int64
OfflineLogicalSessions int64
LogicalOutboxFrames int64
LogicalOutboxBytes int64
PendingPushBytes int64
InboundRPCTasks int64
InboundRPCBytes int64
InboundRPCReadyConnections int64
InboundRPCMaxTasks int64
InboundRPCMaxBytes int64
InboundFrameBytes int64
InboundFrameMaxBytes int64
OutboundTrackedBytes int64
OutboundTrackedMaxBytes int64
OutboundControlBytes int64
OutboundControlMaxBytes int64
OutboundWriteBytes int64
OutboundWriteMaxBytes int64
RPCResultOwners int64
RPCResultReceipts int64
RPCResultReceiptBytes int64
RPCResultSubscribers int64
}
type sessionManagerRuntimeSnapshot struct {
active int64
provisional int64
logical int64
offlineLogical int64
frames int64
bytes int64
pendingBytes int64
}
func (m *SessionManager) runtimeSnapshot() sessionManagerRuntimeSnapshot {
if m == nil {
return sessionManagerRuntimeSnapshot{}
}
// Never hold SessionManager.mu while taking an outbound-state mutex. The
// physical actor can publish/retire a Conn next to an outbox transition, and
// metrics must not add a new cross-component lock order.
m.mu.RLock()
states := make([]*outboundState, 0, len(m.logicalSessions))
result := sessionManagerRuntimeSnapshot{
active: int64(len(m.bySession)),
provisional: int64(len(m.claims)),
logical: int64(len(m.logicalSessions)),
}
if m.pendingBudget != nil {
result.pendingBytes = m.pendingBudget.snapshot()
}
for _, logical := range m.logicalSessions {
if logical == nil {
continue
}
if !logical.offlineAt.IsZero() {
result.offlineLogical++
}
if logical.outbound != nil {
states = append(states, logical.outbound)
}
}
m.mu.RUnlock()
for _, state := range states {
state.mu.Lock()
result.frames += int64(len(state.pending))
result.bytes += int64(state.totalBytes)
state.mu.Unlock()
}
return result
}
type admissionRuntimeSnapshot struct {
connections int64
connectionLimit int64
handshakes int64
handshakeLimit int64
}
func (a *admissionController) runtimeSnapshot() admissionRuntimeSnapshot {
if a == nil {
return admissionRuntimeSnapshot{}
}
a.mu.Lock()
result := admissionRuntimeSnapshot{
connections: int64(a.connections),
connectionLimit: int64(a.maxConnections),
}
a.mu.Unlock()
if a.handshakes != nil {
result.handshakes = int64(len(a.handshakes))
result.handshakeLimit = int64(cap(a.handshakes))
}
return result
}
type inboundRPCRuntimeSnapshot struct {
tasks int64
bytes int64
ready int64
}
func (s *inboundRPCScheduler) runtimeSnapshot() inboundRPCRuntimeSnapshot {
if s == nil {
return inboundRPCRuntimeSnapshot{}
}
s.budgetMu.Lock()
result := inboundRPCRuntimeSnapshot{tasks: int64(s.tasks), bytes: s.bytes}
s.budgetMu.Unlock()
s.readyMu.Lock()
result.ready = int64(s.ready.Len())
s.readyMu.Unlock()
return result
}
// RuntimeSnapshot returns aggregate MTProto ownership and capacity state.
func (s *Server) RuntimeSnapshot() RuntimeSnapshot {
if s == nil {
return RuntimeSnapshot{}
}
sessions := s.conns.runtimeSnapshot()
admission := s.admission.runtimeSnapshot()
inbound := s.rpcScheduler.runtimeSnapshot()
result := RuntimeSnapshot{
RawConnections: admission.connections,
RawConnectionLimit: admission.connectionLimit,
Handshakes: admission.handshakes,
HandshakeLimit: admission.handshakeLimit,
ActiveSessions: sessions.active,
ProvisionalSessions: sessions.provisional,
LogicalSessions: sessions.logical,
OfflineLogicalSessions: sessions.offlineLogical,
LogicalOutboxFrames: sessions.frames,
LogicalOutboxBytes: sessions.bytes,
PendingPushBytes: sessions.pendingBytes,
InboundRPCTasks: inbound.tasks,
InboundRPCBytes: inbound.bytes,
InboundRPCReadyConnections: inbound.ready,
}
if s.rpcScheduler != nil {
result.InboundRPCMaxTasks = int64(s.rpcScheduler.maxTasks)
result.InboundRPCMaxBytes = s.rpcScheduler.maxBytes
}
if s.frameBudget != nil {
result.InboundFrameBytes = s.frameBudget.usedBytes()
result.InboundFrameMaxBytes = s.frameBudget.max
}
if s.outboundTrackedBudget != nil {
result.OutboundTrackedBytes = s.outboundTrackedBudget.snapshot()
result.OutboundTrackedMaxBytes = s.outboundTrackedBudget.maxBytes
}
if s.outboundControlBudget != nil {
result.OutboundControlBytes = s.outboundControlBudget.snapshot()
result.OutboundControlMaxBytes = s.outboundControlBudget.maxBytes
}
if s.outboundScratchPool != nil && s.outboundScratchPool.budget != nil {
result.OutboundWriteBytes = s.outboundScratchPool.snapshot()
result.OutboundWriteMaxBytes = s.outboundScratchPool.budget.maxBytes
}
if s.rpcResults != nil {
result.RPCResultOwners = s.rpcResults.flightLimit.snapshot()
result.RPCResultReceipts = s.rpcResults.completedEntries.snapshot()
result.RPCResultReceiptBytes = s.rpcResults.completedBytes.snapshot()
if s.rpcResults.subscriberBudget != nil {
result.RPCResultSubscribers = s.rpcResults.subscriberBudget.global.snapshot()
}
}
return result
}

View file

@ -0,0 +1,27 @@
package mtprotoedge
import "testing"
func TestRuntimeSnapshotIsNilSafeAndReportsConfiguredLimits(t *testing.T) {
if got := (*Server)(nil).RuntimeSnapshot(); got != (RuntimeSnapshot{}) {
t.Fatalf("nil server snapshot = %#v, want zero", got)
}
if got := (&Server{}).RuntimeSnapshot(); got != (RuntimeSnapshot{}) {
t.Fatalf("partial server snapshot = %#v, want zero", got)
}
server := New(Options{})
snapshot := server.RuntimeSnapshot()
if snapshot.RawConnectionLimit <= 0 || snapshot.HandshakeLimit <= 0 {
t.Fatalf("admission limits not reported: %#v", snapshot)
}
if snapshot.InboundRPCMaxTasks <= 0 || snapshot.InboundRPCMaxBytes <= 0 {
t.Fatalf("inbound RPC limits not reported: %#v", snapshot)
}
if snapshot.InboundFrameMaxBytes <= 0 || snapshot.OutboundTrackedMaxBytes <= 0 || snapshot.OutboundWriteMaxBytes <= 0 {
t.Fatalf("byte limits not reported: %#v", snapshot)
}
if snapshot.RawConnections != 0 || snapshot.ActiveSessions != 0 || snapshot.LogicalOutboxBytes != 0 {
t.Fatalf("fresh server reported live ownership: %#v", snapshot)
}
}

View file

@ -0,0 +1,618 @@
// Package metrics provides a dependency-free Prometheus text exporter for the
// bounded runtime signals emitted by telesrv. It deliberately accepts only a
// small fixed label shape and caps dynamic series so observability cannot become
// an attacker-controlled memory cache.
package metrics
import (
"context"
"errors"
"fmt"
"net/http"
"sort"
"strconv"
"strings"
"sync"
"sync/atomic"
"time"
)
const (
// One additional fixed series reports observations rejected by this cap, so
// the complete exporter remains bounded to 8192 resident series.
defaultMaxSeries = int64(8191)
maxLabelBytes = 96
maxProviderSeries = 1024
)
var durationBuckets = [...]time.Duration{
time.Millisecond,
5 * time.Millisecond,
10 * time.Millisecond,
25 * time.Millisecond,
50 * time.Millisecond,
100 * time.Millisecond,
250 * time.Millisecond,
500 * time.Millisecond,
time.Second,
2 * time.Second,
5 * time.Second,
10 * time.Second,
30 * time.Second,
}
// Label is a bounded Prometheus label attached to a provider sample.
type Label struct {
Name string
Value string
}
// GaugeSample is an identity-free point-in-time value supplied at scrape time.
type GaugeSample struct {
Name string
Labels []Label
Value float64
}
// GaugeProvider is evaluated only during a scrape. Providers must be bounded
// and must not perform unbounded database scans.
type GaugeProvider func() []GaugeSample
type seriesKey struct {
name string
k1 string
v1 string
k2 string
v2 string
k3 string
v3 string
}
func newSeriesKey(name string, labels ...Label) seriesKey {
key := seriesKey{name: sanitizeMetricName(name)}
if len(labels) > 0 {
key.k1, key.v1 = sanitizeLabelName(labels[0].Name), sanitizeLabelValue(labels[0].Value)
}
if len(labels) > 1 {
key.k2, key.v2 = sanitizeLabelName(labels[1].Name), sanitizeLabelValue(labels[1].Value)
}
if len(labels) > 2 {
key.k3, key.v3 = sanitizeLabelName(labels[2].Name), sanitizeLabelValue(labels[2].Value)
}
return key
}
func (k seriesKey) overflow() seriesKey {
if k.k1 != "" {
k.v1 = "overflow"
}
if k.k2 != "" {
k.v2 = "overflow"
}
if k.k3 != "" {
k.v3 = "overflow"
}
return k
}
type counterSeries struct {
key seriesKey
value atomic.Uint64
}
type gaugeSeries struct {
key seriesKey
value atomic.Int64
}
type histogramSeries struct {
key seriesKey
buckets [len(durationBuckets)]atomic.Uint64
count atomic.Uint64
sumNS atomic.Int64
}
func (h *histogramSeries) observe(d time.Duration) {
if d < 0 {
d = 0
}
h.count.Add(1)
h.sumNS.Add(int64(d))
for i, bound := range durationBuckets {
if d <= bound {
h.buckets[i].Add(1)
}
}
}
// Registry implements the mtprotoedge and rpc metric hooks and serves the
// Prometheus text exposition format.
type Registry struct {
maxSeries int64
series atomic.Int64
seriesMu sync.Mutex
dropped atomic.Uint64
counters sync.Map // seriesKey -> *counterSeries
gauges sync.Map // seriesKey -> *gaugeSeries
hist sync.Map // seriesKey -> *histogramSeries
providersMu sync.RWMutex
providers []GaugeProvider
}
// New returns an empty bounded registry.
func New() *Registry {
return &Registry{maxSeries: defaultMaxSeries}
}
// AddGaugeProvider registers a bounded point-in-time provider.
func (r *Registry) AddGaugeProvider(provider GaugeProvider) {
if r == nil || provider == nil {
return
}
r.providersMu.Lock()
r.providers = append(r.providers, provider)
r.providersMu.Unlock()
}
func (r *Registry) counter(key seriesKey) *counterSeries {
if existing, ok := r.counters.Load(key); ok {
return existing.(*counterSeries)
}
r.seriesMu.Lock()
defer r.seriesMu.Unlock()
if existing, ok := r.counters.Load(key); ok {
return existing.(*counterSeries)
}
if r.series.Load() >= r.maxSeries {
r.dropped.Add(1)
return &counterSeries{}
}
created := &counterSeries{key: key}
r.counters.Store(key, created)
r.series.Add(1)
return created
}
func (r *Registry) gauge(key seriesKey) *gaugeSeries {
if existing, ok := r.gauges.Load(key); ok {
return existing.(*gaugeSeries)
}
r.seriesMu.Lock()
defer r.seriesMu.Unlock()
if existing, ok := r.gauges.Load(key); ok {
return existing.(*gaugeSeries)
}
if r.series.Load() >= r.maxSeries {
r.dropped.Add(1)
return &gaugeSeries{}
}
created := &gaugeSeries{key: key}
r.gauges.Store(key, created)
r.series.Add(1)
return created
}
func (r *Registry) histogram(key seriesKey) *histogramSeries {
if existing, ok := r.hist.Load(key); ok {
return existing.(*histogramSeries)
}
r.seriesMu.Lock()
defer r.seriesMu.Unlock()
if existing, ok := r.hist.Load(key); ok {
return existing.(*histogramSeries)
}
if r.series.Load() >= r.maxSeries {
r.dropped.Add(1)
return &histogramSeries{}
}
created := &histogramSeries{key: key}
r.hist.Store(key, created)
r.series.Add(1)
return created
}
func (r *Registry) inc(name string, labels ...Label) {
if r == nil {
return
}
r.counter(newSeriesKey(name, labels...)).value.Add(1)
}
func (r *Registry) add(name string, value uint64, labels ...Label) {
if r == nil || value == 0 {
return
}
r.counter(newSeriesKey(name, labels...)).value.Add(value)
}
func (r *Registry) addGauge(name string, delta int64, labels ...Label) {
if r == nil || delta == 0 {
return
}
r.gauge(newSeriesKey(name, labels...)).value.Add(delta)
}
func (r *Registry) observe(name string, d time.Duration, labels ...Label) {
if r == nil {
return
}
r.histogram(newSeriesKey(name, labels...)).observe(d)
}
// ConnOpened implements mtprotoedge.Metrics.
func (r *Registry) ConnOpened() {
r.inc("telesrv_mtproto_connections_opened_total")
r.addGauge("telesrv_mtproto_connections_active", 1)
}
// ConnClosed implements mtprotoedge.Metrics.
func (r *Registry) ConnClosed() {
r.inc("telesrv_mtproto_connections_closed_total")
r.addGauge("telesrv_mtproto_connections_active", -1)
}
// HandshakeDone implements mtprotoedge.Metrics.
func (r *Registry) HandshakeDone(d time.Duration) {
r.inc("telesrv_mtproto_handshakes_total")
r.observe("telesrv_mtproto_handshake_duration_seconds", d)
}
// RPCHandled implements mtprotoedge.Metrics.
func (r *Registry) RPCHandled(method string, d time.Duration, err error) {
labels := []Label{{Name: "method", Value: method}, {Name: "outcome", Value: errorOutcome(err)}}
r.inc("telesrv_mtproto_rpc_handled_total", labels...)
r.observe("telesrv_mtproto_rpc_duration_seconds", d, labels...)
}
// InboundRPCQueued implements mtprotoedge.Metrics.
func (r *Registry) InboundRPCQueued(method string, length, capacity int) {
r.inc("telesrv_mtproto_inbound_rpc_queued_total", Label{Name: "method", Value: method})
r.add("telesrv_mtproto_inbound_rpc_queue_depth_observed_total", uint64(max(length, 0)), Label{Name: "method", Value: method})
if capacity > 0 && length >= capacity {
r.inc("telesrv_mtproto_inbound_rpc_queue_full_total", Label{Name: "method", Value: method})
}
}
// InboundRPCStarted implements mtprotoedge.Metrics.
func (r *Registry) InboundRPCStarted(method string, queueWait time.Duration) {
r.observe("telesrv_mtproto_inbound_rpc_queue_wait_seconds", queueWait, Label{Name: "method", Value: method})
}
// InboundRPCDropped implements mtprotoedge.Metrics.
func (r *Registry) InboundRPCDropped(method, reason string) {
r.inc("telesrv_mtproto_inbound_rpc_dropped_total", Label{Name: "method", Value: method}, Label{Name: "reason", Value: reason})
}
// OutboundSend implements mtprotoedge.Metrics.
func (r *Registry) OutboundSend(typeID uint32, queueWait time.Duration, bytes int, err error) {
labels := []Label{{Name: "type_id", Value: fmt.Sprintf("%08x", typeID)}, {Name: "outcome", Value: errorOutcome(err)}}
r.inc("telesrv_mtproto_outbound_send_total", labels...)
r.add("telesrv_mtproto_outbound_send_bytes_total", uint64(max(bytes, 0)), labels...)
r.observe("telesrv_mtproto_outbound_queue_wait_seconds", queueWait, labels...)
}
// OutboundResend implements mtprotoedge.Metrics.
func (r *Registry) OutboundResend(count int, err error) {
labels := []Label{{Name: "outcome", Value: errorOutcome(err)}}
r.inc("telesrv_mtproto_outbound_resend_requests_total", labels...)
r.add("telesrv_mtproto_outbound_resent_frames_total", uint64(max(count, 0)), labels...)
}
// OutboundDropped implements mtprotoedge.Metrics.
func (r *Registry) OutboundDropped(reason string) {
r.inc("telesrv_mtproto_outbound_dropped_total", Label{Name: "reason", Value: reason})
}
// OutboundQueueWait implements mtprotoedge.Metrics.
func (r *Registry) OutboundQueueWait(length, capacity int) {
r.inc("telesrv_mtproto_outbound_queue_wait_total")
if capacity > 0 && length >= capacity {
r.inc("telesrv_mtproto_outbound_queue_full_total")
}
}
// RPCResultPrepared implements mtprotoedge.RPCResultMetrics.
func (r *Registry) RPCResultPrepared(method, priority string, innerBytes, wireBytes int, compressed bool) {
encoding := "plain"
if compressed {
encoding = "gzip"
}
labels := []Label{{Name: "method", Value: method}, {Name: "priority", Value: priority}, {Name: "encoding", Value: encoding}}
r.inc("telesrv_mtproto_rpc_result_prepared_total", labels...)
r.add("telesrv_mtproto_rpc_result_inner_bytes_total", uint64(max(innerBytes, 0)), labels...)
r.add("telesrv_mtproto_rpc_result_wire_bytes_total", uint64(max(wireBytes, 0)), labels...)
}
// RPCResultDelivered implements mtprotoedge.RPCResultMetrics.
func (r *Registry) RPCResultDelivered(method string, egressLatency time.Duration, wireBytes int, err error) {
labels := []Label{{Name: "method", Value: method}, {Name: "outcome", Value: errorOutcome(err)}}
r.inc("telesrv_mtproto_rpc_result_delivered_total", labels...)
r.add("telesrv_mtproto_rpc_result_delivered_bytes_total", uint64(max(wireBytes, 0)), labels...)
r.observe("telesrv_mtproto_rpc_result_egress_seconds", egressLatency, labels...)
}
// LogicalOutboxAcknowledged implements mtprotoedge.LogicalOutboxMetrics.
func (r *Registry) LogicalOutboxAcknowledged(bytes int, retainedFor time.Duration, rpcResult bool) {
kind := "service_or_update"
if rpcResult {
kind = "rpc_result"
}
labels := []Label{{Name: "kind", Value: kind}}
r.inc("telesrv_mtproto_logical_outbox_acked_frames_total", labels...)
r.add("telesrv_mtproto_logical_outbox_acked_bytes_total", uint64(max(bytes, 0)), labels...)
r.observe("telesrv_mtproto_logical_outbox_retained_seconds", retainedFor, labels...)
}
// ConnectionIntake implements mtprotoedge.ConnectionIntakeMetrics.
func (r *Registry) ConnectionIntake(stage, outcome string, d time.Duration) {
labels := []Label{{Name: "stage", Value: stage}, {Name: "outcome", Value: outcome}}
r.inc("telesrv_mtproto_connection_intake_total", labels...)
r.observe("telesrv_mtproto_connection_intake_seconds", d, labels...)
}
// MessageSend implements rpc.Metrics.
func (r *Registry) MessageSend(d time.Duration, duplicate bool, err error) {
dup := "false"
if duplicate {
dup = "true"
}
labels := []Label{{Name: "outcome", Value: errorOutcome(err)}, {Name: "duplicate", Value: dup}}
r.inc("telesrv_rpc_message_send_total", labels...)
r.observe("telesrv_rpc_message_send_duration_seconds", d, labels...)
}
// MessageRateLimited implements rpc.Metrics.
func (r *Registry) MessageRateLimited(retryAfterSeconds int) {
r.inc("telesrv_rpc_message_rate_limited_total")
r.add("telesrv_rpc_message_rate_limit_wait_seconds_total", uint64(max(retryAfterSeconds, 0)))
}
// OutboxClaimed implements rpc.Metrics.
func (r *Registry) OutboxClaimed(count int) {
r.add("telesrv_rpc_outbox_claimed_total", uint64(max(count, 0)))
}
// OutboxDelivered implements rpc.Metrics.
func (r *Registry) OutboxDelivered(d time.Duration) {
r.inc("telesrv_rpc_outbox_delivered_total")
r.observe("telesrv_rpc_outbox_delivery_seconds", d)
}
// OutboxFailed implements rpc.Metrics.
func (r *Registry) OutboxFailed(err error) {
r.inc("telesrv_rpc_outbox_failed_total", Label{Name: "outcome", Value: errorOutcome(err)})
}
// ServeHTTP writes Prometheus text format.
func (r *Registry) ServeHTTP(w http.ResponseWriter, _ *http.Request) {
w.Header().Set("Content-Type", "text/plain; version=0.0.4; charset=utf-8")
var counters []*counterSeries
r.counters.Range(func(_, value any) bool {
counters = append(counters, value.(*counterSeries))
return true
})
var gauges []*gaugeSeries
r.gauges.Range(func(_, value any) bool {
gauges = append(gauges, value.(*gaugeSeries))
return true
})
var histograms []*histogramSeries
r.hist.Range(func(_, value any) bool {
histograms = append(histograms, value.(*histogramSeries))
return true
})
providerSamples := r.providerSamples()
sort.Slice(counters, func(i, j int) bool { return lessSeries(counters[i].key, counters[j].key) })
sort.Slice(gauges, func(i, j int) bool { return lessSeries(gauges[i].key, gauges[j].key) })
sort.Slice(histograms, func(i, j int) bool { return lessSeries(histograms[i].key, histograms[j].key) })
sort.Slice(providerSamples, func(i, j int) bool {
if providerSamples[i].Name != providerSamples[j].Name {
return providerSamples[i].Name < providerSamples[j].Name
}
return labelsString(providerSamples[i].Labels) < labelsString(providerSamples[j].Labels)
})
var out strings.Builder
fmt.Fprintln(&out, "# TYPE telesrv_metrics_dropped_observations_total counter")
fmt.Fprintf(&out, "telesrv_metrics_dropped_observations_total %d\n", r.dropped.Load())
writeCounterSeries(&out, counters)
writeGaugeSeries(&out, gauges, providerSamples)
writeHistogramSeries(&out, histograms)
_, _ = w.Write([]byte(out.String()))
}
func (r *Registry) providerSamples() (samples []GaugeSample) {
r.providersMu.RLock()
providers := append([]GaugeProvider(nil), r.providers...)
r.providersMu.RUnlock()
for _, provider := range providers {
func() {
defer func() { _ = recover() }()
for _, sample := range provider() {
if len(samples) >= maxProviderSeries {
r.dropped.Add(1)
return
}
sample.Name = sanitizeMetricName(sample.Name)
if len(sample.Labels) > 3 {
sample.Labels = sample.Labels[:3]
}
for i := range sample.Labels {
sample.Labels[i].Name = sanitizeLabelName(sample.Labels[i].Name)
sample.Labels[i].Value = sanitizeLabelValue(sample.Labels[i].Value)
}
samples = append(samples, sample)
}
}()
}
return samples
}
func writeCounterSeries(out *strings.Builder, series []*counterSeries) {
last := ""
for _, item := range series {
if item.key.name != last {
fmt.Fprintf(out, "# TYPE %s counter\n", item.key.name)
last = item.key.name
}
writeSample(out, item.key.name, keyLabels(item.key), float64(item.value.Load()))
}
}
func writeGaugeSeries(out *strings.Builder, series []*gaugeSeries, provider []GaugeSample) {
type sample struct {
name string
labels []Label
value float64
}
all := make([]sample, 0, len(series)+len(provider))
for _, item := range series {
all = append(all, sample{name: item.key.name, labels: keyLabels(item.key), value: float64(item.value.Load())})
}
for _, item := range provider {
all = append(all, sample{name: item.Name, labels: item.Labels, value: item.Value})
}
sort.Slice(all, func(i, j int) bool {
if all[i].name != all[j].name {
return all[i].name < all[j].name
}
return labelsString(all[i].labels) < labelsString(all[j].labels)
})
last := ""
for _, item := range all {
if item.name != last {
fmt.Fprintf(out, "# TYPE %s gauge\n", item.name)
last = item.name
}
writeSample(out, item.name, item.labels, item.value)
}
}
func writeHistogramSeries(out *strings.Builder, series []*histogramSeries) {
last := ""
for _, item := range series {
if item.key.name != last {
fmt.Fprintf(out, "# TYPE %s histogram\n", item.key.name)
last = item.key.name
}
labels := keyLabels(item.key)
for i, bound := range durationBuckets {
bucketLabels := append(append([]Label(nil), labels...), Label{Name: "le", Value: strconv.FormatFloat(bound.Seconds(), 'g', -1, 64)})
writeSample(out, item.key.name+"_bucket", bucketLabels, float64(item.buckets[i].Load()))
}
writeSample(out, item.key.name+"_bucket", append(append([]Label(nil), labels...), Label{Name: "le", Value: "+Inf"}), float64(item.count.Load()))
writeSample(out, item.key.name+"_sum", labels, time.Duration(item.sumNS.Load()).Seconds())
writeSample(out, item.key.name+"_count", labels, float64(item.count.Load()))
}
}
func writeSample(out *strings.Builder, name string, labels []Label, value float64) {
out.WriteString(name)
if len(labels) > 0 {
out.WriteByte('{')
for i, label := range labels {
if i > 0 {
out.WriteByte(',')
}
out.WriteString(label.Name)
out.WriteString("=\"")
out.WriteString(escapeLabel(label.Value))
out.WriteByte('"')
}
out.WriteByte('}')
}
out.WriteByte(' ')
out.WriteString(strconv.FormatFloat(value, 'g', -1, 64))
out.WriteByte('\n')
}
func keyLabels(key seriesKey) []Label {
labels := make([]Label, 0, 3)
if key.k1 != "" {
labels = append(labels, Label{Name: key.k1, Value: key.v1})
}
if key.k2 != "" {
labels = append(labels, Label{Name: key.k2, Value: key.v2})
}
if key.k3 != "" {
labels = append(labels, Label{Name: key.k3, Value: key.v3})
}
return labels
}
func lessSeries(a, b seriesKey) bool {
if a.name != b.name {
return a.name < b.name
}
return a.k1+a.v1+a.k2+a.v2+a.k3+a.v3 < b.k1+b.v1+b.k2+b.v2+b.k3+b.v3
}
func labelsString(labels []Label) string {
var b strings.Builder
for _, label := range labels {
b.WriteString(label.Name)
b.WriteByte('=')
b.WriteString(label.Value)
b.WriteByte(',')
}
return b.String()
}
func errorOutcome(err error) string {
if err == nil {
return "ok"
}
if errors.Is(err, context.Canceled) {
return "canceled"
}
if errors.Is(err, context.DeadlineExceeded) {
return "timeout"
}
message := strings.ToUpper(err.Error())
switch {
case strings.Contains(message, "FLOOD_WAIT"):
return "flood_wait"
case strings.Contains(message, "WORKER_BUSY") || strings.Contains(message, "BUDGET") || strings.Contains(message, "CAPACITY"):
return "edge_overload"
default:
return "error"
}
}
func sanitizeMetricName(value string) string {
if value == "" {
return "telesrv_invalid_metric"
}
var b strings.Builder
for i, r := range value {
valid := r == '_' || r == ':' || r >= 'a' && r <= 'z' || r >= 'A' && r <= 'Z' || i > 0 && r >= '0' && r <= '9'
if valid {
b.WriteRune(r)
} else {
b.WriteByte('_')
}
}
return b.String()
}
func sanitizeLabelName(value string) string {
return strings.ReplaceAll(sanitizeMetricName(value), ":", "_")
}
func sanitizeLabelValue(value string) string {
if value == "" {
return "unknown"
}
if len(value) > maxLabelBytes {
return value[:maxLabelBytes]
}
return value
}
func escapeLabel(value string) string {
value = strings.ReplaceAll(value, "\\", "\\\\")
value = strings.ReplaceAll(value, "\n", "\\n")
return strings.ReplaceAll(value, "\"", "\\\"")
}

View file

@ -0,0 +1,83 @@
package metrics
import (
"errors"
"net/http/httptest"
"strings"
"testing"
"time"
"telesrv/internal/mtprotoedge"
"telesrv/internal/rpc"
)
var (
_ mtprotoedge.Metrics = (*Registry)(nil)
_ mtprotoedge.RPCResultMetrics = (*Registry)(nil)
_ mtprotoedge.LogicalOutboxMetrics = (*Registry)(nil)
_ mtprotoedge.ConnectionIntakeMetrics = (*Registry)(nil)
_ rpc.Metrics = (*Registry)(nil)
)
func TestRegistryExportsBoundedAggregateMetrics(t *testing.T) {
registry := New()
registry.maxSeries = 2
registry.RPCHandled("help.getConfig", 5*time.Millisecond, nil)
registry.RPCHandled("users.getUsers", time.Second, errors.New("secret auth_key_id=deadbeef session=123"))
recorder := httptest.NewRecorder()
registry.ServeHTTP(recorder, httptest.NewRequest("GET", "/metrics", nil))
body := recorder.Body.String()
if recorder.Code != 200 {
t.Fatalf("status = %d, want 200", recorder.Code)
}
if !strings.Contains(body, "telesrv_mtproto_rpc_handled_total") || !strings.Contains(body, "telesrv_mtproto_rpc_duration_seconds_bucket") {
t.Fatalf("expected RPC counter and histogram, got:\n%s", body)
}
if strings.Contains(body, "deadbeef") || strings.Contains(body, "session=123") {
t.Fatalf("raw error identity leaked into metrics:\n%s", body)
}
if got := registry.series.Load(); got != registry.maxSeries {
t.Fatalf("resident dynamic series = %d, want cap %d", got, registry.maxSeries)
}
if got := registry.dropped.Load(); got == 0 {
t.Fatal("series overflow was not reported")
}
if !strings.Contains(body, "telesrv_metrics_dropped_observations_total 2") {
t.Fatalf("overflow counter missing from:\n%s", body)
}
}
func TestRegistrySanitizesAndBoundsProviderSamples(t *testing.T) {
registry := New()
registry.AddGaugeProvider(func() []GaugeSample {
return []GaugeSample{{
Name: "9 invalid metric",
Labels: []Label{{Name: "bad:label", Value: "quoted\"\nvalue"}},
Value: 3,
}}
})
recorder := httptest.NewRecorder()
registry.ServeHTTP(recorder, httptest.NewRequest("GET", "/metrics", nil))
body := recorder.Body.String()
if !strings.Contains(body, `__invalid_metric{bad_label="quoted\"\nvalue"} 3`) {
t.Fatalf("provider sample was not safely sanitized:\n%s", body)
}
}
func TestErrorOutcomeHasFixedCardinality(t *testing.T) {
tests := []struct {
err error
want string
}{
{nil, "ok"},
{errors.New("FLOOD_WAIT_1 for phone 123"), "flood_wait"},
{errors.New("global capacity exceeded for auth key"), "edge_overload"},
{errors.New("arbitrary user-controlled failure"), "error"},
}
for _, test := range tests {
if got := errorOutcome(test.err); got != test.want {
t.Errorf("errorOutcome(%v) = %q, want %q", test.err, got, test.want)
}
}
}

View file

@ -67,6 +67,13 @@ func (r *Router) checkSendRateLimit(ctx context.Context, userID int64, cost int)
}
allowed, retryAfter, err := r.deps.Limiter.AllowN(ctx, sendRateLimitKeyPrefix+strconv.FormatInt(userID, 10), cost, limit, window)
if err != nil {
r.log.Warn("message send rate limiter failed",
append(r.contextLogFields(ctx),
zap.Error(err),
zap.Int("cost", cost),
zap.Int("limit", limit),
zap.Duration("window", window),
)...)
return internalErr()
}
if allowed {

View file

@ -66,8 +66,12 @@ func TestDispatchUnwrapsWrappers(t *testing.T) {
if cfg.ThisDC != dc {
t.Fatalf("ThisDC = %d, want %d", cfg.ThisDC, dc)
}
if len(cfg.DCOptions) != 0 {
t.Fatalf("DCOptions = %+v, want empty (client uses pinned static address)", cfg.DCOptions)
if len(cfg.DCOptions) != 1 {
t.Fatalf("DCOptions = %+v, want one reconnect route", cfg.DCOptions)
}
option := cfg.DCOptions[0]
if option.ID != dc || option.IPAddress != ip || option.Port != port || option.Ipv6 || option.MediaOnly || option.CDN {
t.Fatalf("DCOptions[0] = %+v, want primary dc=%d at %s:%d", option, dc, ip, port)
}
}

View file

@ -5,6 +5,7 @@ import (
"crypto/sha256"
"github.com/iamxvbaba/td/tg"
"go.uber.org/zap"
"telesrv/internal/domain"
)
@ -46,6 +47,11 @@ func (r *Router) lookupOutgoingReplay(ctx context.Context, userID int64, peer do
IdempotencyFingerprint: fingerprint,
})
if err != nil {
r.log.Warn("private send replay lookup failed",
append(r.contextLogFields(ctx),
zap.Error(err),
zap.Int64("recipient_user_id", peer.ID),
)...)
return outgoingReplayLookup{checked: true}, messageSendErr(err)
}
return outgoingReplayLookup{private: res, found: found, checked: true}, nil

View file

@ -58,21 +58,32 @@ func (l *RateLimiter) AllowN(ctx context.Context, key string, cost, limit int, w
if err != nil {
return false, 0, fmt.Errorf("redis increment rate limit: %w", err)
}
items, ok := value.([]interface{})
if !ok || len(items) != 2 {
return false, 0, fmt.Errorf("redis increment rate limit: unexpected result %T", value)
}
count, countOK := items[0].(int64)
ttlMillis, ttlOK := items[1].(int64)
if !countOK || !ttlOK || ttlMillis <= 0 {
return false, 0, fmt.Errorf("redis increment rate limit: invalid result %#v", items)
count, ttlMillis, err := decodeRateLimitIncrementResult(value)
if err != nil {
return false, 0, err
}
if count <= int64(limit) {
return true, 0, nil
}
// Redis PTTL returns 0 when less than one millisecond remains. That is a
// valid fixed-window boundary, not a corrupt result. Round it up to the
// smallest protocol-safe FLOOD_WAIT instead of leaking a transient 500.
retry := (ttlMillis + 999) / 1000
if retry <= 0 {
retry = 1
}
return false, int(retry), nil
}
func decodeRateLimitIncrementResult(value any) (count int64, ttlMillis int64, err error) {
items, ok := value.([]interface{})
if !ok || len(items) != 2 {
return 0, 0, fmt.Errorf("redis increment rate limit: unexpected result %T", value)
}
count, countOK := items[0].(int64)
ttlMillis, ttlOK := items[1].(int64)
if !countOK || !ttlOK || count <= 0 || ttlMillis < 0 {
return 0, 0, fmt.Errorf("redis increment rate limit: invalid result %#v", items)
}
return count, ttlMillis, nil
}

View file

@ -0,0 +1,32 @@
package redisstore
import "testing"
func TestDecodeRateLimitIncrementResultAcceptsPTTLBoundary(t *testing.T) {
count, ttlMillis, err := decodeRateLimitIncrementResult([]interface{}{int64(7), int64(0)})
if err != nil {
t.Fatalf("decode zero PTTL: %v", err)
}
if count != 7 || ttlMillis != 0 {
t.Fatalf("decoded count=%d ttl=%d, want 7/0", count, ttlMillis)
}
}
func TestDecodeRateLimitIncrementResultRejectsInvalidShape(t *testing.T) {
tests := []struct {
name string
value any
}{
{name: "wrong type", value: "7,1"},
{name: "wrong length", value: []interface{}{int64(7)}},
{name: "zero count", value: []interface{}{int64(0), int64(1)}},
{name: "negative ttl", value: []interface{}{int64(7), int64(-1)}},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
if _, _, err := decodeRateLimitIncrementResult(test.value); err == nil {
t.Fatal("expected decode error")
}
})
}
}