owpengram-server/internal/loadharness/provision.go

267 lines
7.9 KiB
Go

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)