| 1 | package identity |
| 2 | |
| 3 | import ( |
| 4 | "encoding/json" |
| 5 | "errors" |
| 6 | "fmt" |
| 7 | "math" |
| 8 | "net" |
| 9 | "os" |
| 10 | "path/filepath" |
| 11 | "strings" |
| 12 | |
| 13 | "github.com/gosuda/portal-tunnel/v2/types" |
| 14 | "github.com/gosuda/portal-tunnel/v2/utils" |
| 15 | ) |
| 16 | |
| 17 | func NormalizeIdentity(identity types.Identity) (types.Identity, error) { |
| 18 | normalized := identity.Copy() |
| 19 | |
| 20 | name, err := utils.NormalizeDNSLabel(identity.Name) |
| 21 | if err != nil { |
| 22 | return types.Identity{}, err |
| 23 | } |
| 24 | address, err := NormalizeEVMAddress(identity.Address) |
| 25 | if err != nil { |
| 26 | return types.Identity{}, err |
| 27 | } |
| 28 | |
| 29 | normalized.Name = name |
| 30 | normalized.Address = address |
| 31 | return normalized, nil |
| 32 | } |
| 33 | |
| 34 | func NormalizeRelayDescriptor(desc types.RelayDescriptor) (types.RelayDescriptor, error) { |
| 35 | desc.Address = strings.TrimSpace(desc.Address) |
| 36 | desc.Version = strings.TrimSpace(desc.Version) |
| 37 | desc.APIHTTPSAddr = strings.TrimSpace(desc.APIHTTPSAddr) |
| 38 | desc.WireGuardPublicKey = strings.TrimSpace(desc.WireGuardPublicKey) |
| 39 | if desc.Version == "" { |
| 40 | desc.Version = types.DiscoveryVersion |
| 41 | } |
| 42 | if !desc.IssuedAt.IsZero() { |
| 43 | desc.IssuedAt = desc.IssuedAt.UTC() |
| 44 | } |
| 45 | if !desc.ExpiresAt.IsZero() { |
| 46 | desc.ExpiresAt = desc.ExpiresAt.UTC() |
| 47 | } |
| 48 | |
| 49 | if desc.APIHTTPSAddr != "" { |
| 50 | normalized, err := utils.NormalizeRelayURL(desc.APIHTTPSAddr) |
| 51 | if err != nil { |
| 52 | return types.RelayDescriptor{}, fmt.Errorf("normalize api https addr: %w", err) |
| 53 | } |
| 54 | desc.APIHTTPSAddr = normalized |
| 55 | } |
| 56 | if desc.Address != "" { |
| 57 | normalized, err := NormalizeEVMAddress(desc.Address) |
| 58 | if err != nil { |
| 59 | return types.RelayDescriptor{}, fmt.Errorf("normalize address: %w", err) |
| 60 | } |
| 61 | desc.Address = normalized |
| 62 | } |
| 63 | if desc.WireGuardPublicKey != "" { |
| 64 | if err := ValidateWireGuardPublicKey(desc.WireGuardPublicKey); err != nil { |
| 65 | return types.RelayDescriptor{}, err |
| 66 | } |
| 67 | } |
| 68 | if desc.WireGuardPort < 0 || desc.WireGuardPort > 65535 { |
| 69 | return types.RelayDescriptor{}, errors.New("wireguard_port is invalid") |
| 70 | } |
| 71 | if desc.ActiveConnections < 0 { |
| 72 | return types.RelayDescriptor{}, errors.New("active_connections is invalid") |
| 73 | } |
| 74 | if desc.TCPBPS < 0 || math.IsNaN(desc.TCPBPS) || math.IsInf(desc.TCPBPS, 0) { |
| 75 | return types.RelayDescriptor{}, errors.New("tcp_bps is invalid") |
| 76 | } |
| 77 | |
| 78 | switch { |
| 79 | case desc.Address == "": |
| 80 | return types.RelayDescriptor{}, errors.New("address is required") |
| 81 | case desc.Version != types.DiscoveryVersion: |
| 82 | return types.RelayDescriptor{}, fmt.Errorf("unsupported relay descriptor version %q", desc.Version) |
| 83 | case desc.APIHTTPSAddr == "": |
| 84 | return types.RelayDescriptor{}, errors.New("api_https_addr is required") |
| 85 | case desc.SupportsOverlay && desc.WireGuardPublicKey == "": |
| 86 | return types.RelayDescriptor{}, errors.New("wireguard_public_key is required when supports_overlay is set") |
| 87 | case desc.SupportsOverlay && desc.WireGuardPort == 0: |
| 88 | return types.RelayDescriptor{}, errors.New("wireguard_port is required when supports_overlay is set") |
| 89 | case !desc.SupportsOverlay && (desc.WireGuardPublicKey != "" || desc.WireGuardPort != 0): |
| 90 | return types.RelayDescriptor{}, errors.New("supports_overlay is required when wireguard metadata is set") |
| 91 | case desc.ExpiresAt.IsZero(): |
| 92 | return types.RelayDescriptor{}, errors.New("expires_at is required") |
| 93 | case desc.IssuedAt.After(desc.ExpiresAt): |
| 94 | return types.RelayDescriptor{}, errors.New("issued_at must be before expires_at") |
| 95 | } |
| 96 | |
| 97 | return desc, nil |
| 98 | } |
| 99 | |
| 100 | func RelayWireGuardEndpoint(desc types.RelayDescriptor) (string, error) { |
| 101 | host := utils.PortalRootHost(desc.APIHTTPSAddr) |
| 102 | if host == "" { |
| 103 | return "", errors.New("api_https_addr host is required") |
| 104 | } |
| 105 | if desc.WireGuardPort <= 0 || desc.WireGuardPort > 65535 { |
| 106 | return "", errors.New("wireguard_port is invalid") |
| 107 | } |
| 108 | return net.JoinHostPort(host, fmt.Sprintf("%d", desc.WireGuardPort)), nil |
| 109 | } |
| 110 | |
| 111 | func ResolveRelayStateDir(path string) string { |
| 112 | trimmed := strings.TrimSpace(path) |
| 113 | if trimmed == "" { |
| 114 | return "" |
| 115 | } |
| 116 | switch strings.ToLower(filepath.Base(trimmed)) { |
| 117 | case types.RelayIdentityFilename, types.RelayPolicyFilename: |
| 118 | return filepath.Dir(trimmed) |
| 119 | default: |
| 120 | return trimmed |
| 121 | } |
| 122 | } |
| 123 | |
| 124 | func resolveRelayIdentityPath(path string) string { |
| 125 | stateDir := ResolveRelayStateDir(path) |
| 126 | if stateDir == "" { |
| 127 | return "" |
| 128 | } |
| 129 | return filepath.Join(stateDir, types.RelayIdentityFilename) |
| 130 | } |
| 131 | |
| 132 | func ResolveRelayPolicyPath(path string) string { |
| 133 | stateDir := ResolveRelayStateDir(path) |
| 134 | if stateDir == "" { |
| 135 | return "" |
| 136 | } |
| 137 | return filepath.Join(stateDir, types.RelayPolicyFilename) |
| 138 | } |
| 139 | |
| 140 | func normalizeStoredIdentity(identity types.Identity) (types.Identity, error) { |
| 141 | normalized := identity.Copy() |
| 142 | normalized.Name = strings.TrimSpace(normalized.Name) |
| 143 | normalized.Address = strings.TrimSpace(normalized.Address) |
| 144 | normalized.PublicKey = strings.TrimSpace(normalized.PublicKey) |
| 145 | normalized.PrivateKey = strings.TrimSpace(normalized.PrivateKey) |
| 146 | normalized.Mnemonic = normalizeMnemonic(normalized.Mnemonic) |
| 147 | normalized.DerivationPath = strings.TrimSpace(normalized.DerivationPath) |
| 148 | normalized.TokenSecret = strings.TrimSpace(normalized.TokenSecret) |
| 149 | |
| 150 | if normalized.Mnemonic != "" { |
| 151 | privateKey, derivationPath, err := deriveSecp256k1PrivateKeyFromMnemonic(normalized.Mnemonic, normalized.DerivationPath) |
| 152 | if err != nil { |
| 153 | return types.Identity{}, err |
| 154 | } |
| 155 | normalized.DerivationPath = derivationPath |
| 156 | if normalized.PrivateKey == "" { |
| 157 | normalized.PrivateKey = privateKey |
| 158 | } else if !strings.EqualFold(utils.TrimHexPrefix(normalized.PrivateKey), privateKey) { |
| 159 | return types.Identity{}, errors.New("identity private key does not match mnemonic") |
| 160 | } |
| 161 | } else if normalized.DerivationPath != "" { |
| 162 | return types.Identity{}, errors.New("identity derivation_path requires mnemonic") |
| 163 | } |
| 164 | |
| 165 | switch { |
| 166 | case normalized.PrivateKey != "": |
| 167 | resolved, err := ResolveSecp256k1Identity(normalized.PrivateKey) |
| 168 | if err != nil { |
| 169 | return types.Identity{}, err |
| 170 | } |
| 171 | if normalized.PublicKey != "" && !strings.EqualFold(utils.TrimHexPrefix(normalized.PublicKey), resolved.PublicKey) { |
| 172 | return types.Identity{}, errors.New("identity public key does not match private key") |
| 173 | } |
| 174 | if normalized.Address != "" && !strings.EqualFold(normalized.Address, resolved.Address) { |
| 175 | return types.Identity{}, errors.New("identity address does not match private key") |
| 176 | } |
| 177 | normalized.Address = resolved.Address |
| 178 | normalized.PublicKey = resolved.PublicKey |
| 179 | normalized.PrivateKey = resolved.PrivateKey |
| 180 | case normalized.PublicKey != "": |
| 181 | address, err := AddressFromCompressedPublicKeyHex(normalized.PublicKey) |
| 182 | if err != nil { |
| 183 | return types.Identity{}, err |
| 184 | } |
| 185 | normalized.PublicKey = strings.ToLower(utils.TrimHexPrefix(normalized.PublicKey)) |
| 186 | if normalized.Address == "" { |
| 187 | normalized.Address = address |
| 188 | break |
| 189 | } |
| 190 | if !strings.EqualFold(normalized.Address, address) { |
| 191 | return types.Identity{}, errors.New("identity address does not match public key") |
| 192 | } |
| 193 | normalized.Address = address |
| 194 | case normalized.Address != "": |
| 195 | address, err := NormalizeEVMAddress(normalized.Address) |
| 196 | if err != nil { |
| 197 | return types.Identity{}, err |
| 198 | } |
| 199 | normalized.Address = address |
| 200 | } |
| 201 | return normalized, nil |
| 202 | } |
| 203 | |
| 204 | func normalizeStoredRelayIdentity(identity types.RelayIdentity) (types.RelayIdentity, error) { |
| 205 | normalized := identity.Copy() |
| 206 | baseIdentity, err := normalizeStoredIdentity(normalized.Identity) |
| 207 | if err != nil { |
| 208 | return types.RelayIdentity{}, err |
| 209 | } |
| 210 | normalized.Identity = baseIdentity |
| 211 | normalized.WireGuardPublicKey = strings.TrimSpace(normalized.WireGuardPublicKey) |
| 212 | normalized.WireGuardPrivateKey = strings.TrimSpace(normalized.WireGuardPrivateKey) |
| 213 | normalized.EncryptedClientHelloSeed = strings.TrimSpace(normalized.EncryptedClientHelloSeed) |
| 214 | |
| 215 | switch { |
| 216 | case normalized.WireGuardPrivateKey != "": |
| 217 | privateKey, err := NormalizeWireGuardPrivateKey(normalized.WireGuardPrivateKey) |
| 218 | if err != nil { |
| 219 | return types.RelayIdentity{}, fmt.Errorf("normalize wireguard private key: %w", err) |
| 220 | } |
| 221 | publicKey, err := WireGuardPublicKeyFromPrivate(privateKey) |
| 222 | if err != nil { |
| 223 | return types.RelayIdentity{}, fmt.Errorf("derive wireguard public key: %w", err) |
| 224 | } |
| 225 | if configuredPublicKey := strings.TrimSpace(normalized.WireGuardPublicKey); configuredPublicKey != "" { |
| 226 | if err := ValidateWireGuardPublicKey(configuredPublicKey); err != nil { |
| 227 | return types.RelayIdentity{}, err |
| 228 | } |
| 229 | if configuredPublicKey != publicKey { |
| 230 | return types.RelayIdentity{}, errors.New("identity wireguard public key does not match private key") |
| 231 | } |
| 232 | } |
| 233 | normalized.WireGuardPrivateKey = privateKey |
| 234 | normalized.WireGuardPublicKey = publicKey |
| 235 | case normalized.WireGuardPublicKey != "": |
| 236 | if err := ValidateWireGuardPublicKey(normalized.WireGuardPublicKey); err != nil { |
| 237 | return types.RelayIdentity{}, err |
| 238 | } |
| 239 | } |
| 240 | |
| 241 | return normalized, nil |
| 242 | } |
| 243 | |
| 244 | type storedIdentity struct { |
| 245 | Name string `json:"name,omitempty"` |
| 246 | Address string `json:"address,omitempty"` |
| 247 | PublicKey string `json:"public_key,omitempty"` |
| 248 | PrivateKey string `json:"private_key,omitempty"` |
| 249 | Mnemonic string `json:"mnemonic,omitempty"` |
| 250 | DerivationPath string `json:"derivation_path,omitempty"` |
| 251 | TokenSecret string `json:"token_secret,omitempty"` |
| 252 | } |
| 253 | |
| 254 | type storedRelayIdentity struct { |
| 255 | storedIdentity |
| 256 | WireGuardPublicKey string `json:"wireguard_public_key,omitempty"` |
| 257 | WireGuardPrivateKey string `json:"wireguard_private_key,omitempty"` |
| 258 | EncryptedClientHelloSeed string `json:"encrypted_client_hello_seed,omitempty"` |
| 259 | } |
| 260 | |
| 261 | func storedIdentityFromIdentity(identity types.Identity) storedIdentity { |
| 262 | privateKey := identity.PrivateKey |
| 263 | if strings.TrimSpace(identity.Mnemonic) != "" { |
| 264 | privateKey = "" |
| 265 | } |
| 266 | return storedIdentity{ |
| 267 | Name: identity.Name, |
| 268 | Address: identity.Address, |
| 269 | PublicKey: identity.PublicKey, |
| 270 | PrivateKey: privateKey, |
| 271 | Mnemonic: identity.Mnemonic, |
| 272 | DerivationPath: identity.DerivationPath, |
| 273 | TokenSecret: identity.TokenSecret, |
| 274 | } |
| 275 | } |
| 276 | |
| 277 | func saveIdentity(path string, identity types.Identity) error { |
| 278 | path = strings.TrimSpace(path) |
| 279 | if path == "" { |
| 280 | return errors.New("identity path is required") |
| 281 | } |
| 282 | normalized, err := normalizeStoredIdentity(identity) |
| 283 | if err != nil { |
| 284 | return err |
| 285 | } |
| 286 | normalized, err = ensureTokenSecret(normalized) |
| 287 | if err != nil { |
| 288 | return err |
| 289 | } |
| 290 | if err := utils.WriteJSONFile(path, storedIdentityFromIdentity(normalized), 0o600); err != nil { |
| 291 | return fmt.Errorf("write identity file: %w", err) |
| 292 | } |
| 293 | return nil |
| 294 | } |
| 295 | |
| 296 | func saveRelayIdentity(path string, identity types.RelayIdentity) error { |
| 297 | path = resolveRelayIdentityPath(path) |
| 298 | if path == "" { |
| 299 | return errors.New("identity path is required") |
| 300 | } |
| 301 | normalized, err := normalizeStoredRelayIdentity(identity) |
| 302 | if err != nil { |
| 303 | return err |
| 304 | } |
| 305 | baseIdentity, err := ensureTokenSecret(normalized.Identity) |
| 306 | if err != nil { |
| 307 | return err |
| 308 | } |
| 309 | normalized.Identity = baseIdentity |
| 310 | storedBaseIdentity := storedIdentityFromIdentity(normalized.Identity) |
| 311 | if err := utils.WriteJSONFile(path, storedRelayIdentity{ |
| 312 | storedIdentity: storedBaseIdentity, |
| 313 | WireGuardPublicKey: normalized.WireGuardPublicKey, |
| 314 | WireGuardPrivateKey: normalized.WireGuardPrivateKey, |
| 315 | EncryptedClientHelloSeed: normalized.EncryptedClientHelloSeed, |
| 316 | }, 0o600); err != nil { |
| 317 | return fmt.Errorf("write identity file: %w", err) |
| 318 | } |
| 319 | return nil |
| 320 | } |
| 321 | |
| 322 | func loadIdentity(path string) (types.Identity, error) { |
| 323 | path = strings.TrimSpace(path) |
| 324 | if path == "" { |
| 325 | return types.Identity{}, errors.New("identity path is required") |
| 326 | } |
| 327 | var payload storedIdentity |
| 328 | if err := utils.ReadJSONFile(path, &payload); err != nil { |
| 329 | return types.Identity{}, fmt.Errorf("read identity file: %w", err) |
| 330 | } |
| 331 | return normalizeStoredIdentity(types.Identity{ |
| 332 | Name: payload.Name, |
| 333 | Address: payload.Address, |
| 334 | PublicKey: payload.PublicKey, |
| 335 | PrivateKey: payload.PrivateKey, |
| 336 | Mnemonic: payload.Mnemonic, |
| 337 | DerivationPath: payload.DerivationPath, |
| 338 | TokenSecret: payload.TokenSecret, |
| 339 | }) |
| 340 | } |
| 341 | |
| 342 | func loadRelayIdentity(path string) (types.RelayIdentity, error) { |
| 343 | path = resolveRelayIdentityPath(path) |
| 344 | if path == "" { |
| 345 | return types.RelayIdentity{}, errors.New("identity path is required") |
| 346 | } |
| 347 | var payload storedRelayIdentity |
| 348 | if err := utils.ReadJSONFile(path, &payload); err != nil { |
| 349 | return types.RelayIdentity{}, fmt.Errorf("read identity file: %w", err) |
| 350 | } |
| 351 | return normalizeStoredRelayIdentity(types.RelayIdentity{ |
| 352 | Identity: types.Identity{ |
| 353 | Name: payload.Name, |
| 354 | Address: payload.Address, |
| 355 | PublicKey: payload.PublicKey, |
| 356 | PrivateKey: payload.PrivateKey, |
| 357 | Mnemonic: payload.Mnemonic, |
| 358 | DerivationPath: payload.DerivationPath, |
| 359 | TokenSecret: payload.TokenSecret, |
| 360 | }, |
| 361 | WireGuardPublicKey: payload.WireGuardPublicKey, |
| 362 | WireGuardPrivateKey: payload.WireGuardPrivateKey, |
| 363 | EncryptedClientHelloSeed: payload.EncryptedClientHelloSeed, |
| 364 | }) |
| 365 | } |
| 366 | |
| 367 | func parseIdentityJSON(raw string) (types.Identity, error) { |
| 368 | raw = strings.TrimSpace(raw) |
| 369 | if raw == "" { |
| 370 | return types.Identity{}, errors.New("identity json is required") |
| 371 | } |
| 372 | |
| 373 | var payload storedIdentity |
| 374 | if err := json.Unmarshal([]byte(raw), &payload); err != nil { |
| 375 | return types.Identity{}, fmt.Errorf("decode identity json: %w", err) |
| 376 | } |
| 377 | return normalizeStoredIdentity(types.Identity{ |
| 378 | Name: payload.Name, |
| 379 | Address: payload.Address, |
| 380 | PublicKey: payload.PublicKey, |
| 381 | PrivateKey: payload.PrivateKey, |
| 382 | Mnemonic: payload.Mnemonic, |
| 383 | DerivationPath: payload.DerivationPath, |
| 384 | TokenSecret: payload.TokenSecret, |
| 385 | }) |
| 386 | } |
| 387 | |
| 388 | func loadOrCreateIdentity(path string, identity types.Identity) (types.Identity, bool, error) { |
| 389 | path = strings.TrimSpace(path) |
| 390 | if path == "" { |
| 391 | return types.Identity{}, false, errors.New("identity path is required") |
| 392 | } |
| 393 | |
| 394 | stored, err := loadIdentity(path) |
| 395 | switch { |
| 396 | case err == nil: |
| 397 | if name := strings.TrimSpace(identity.Name); name != "" { |
| 398 | stored.Name = name |
| 399 | } |
| 400 | if address := strings.TrimSpace(identity.Address); address != "" { |
| 401 | stored.Address = address |
| 402 | } |
| 403 | if publicKey := strings.TrimSpace(identity.PublicKey); publicKey != "" { |
| 404 | stored.PublicKey = publicKey |
| 405 | } |
| 406 | if privateKey := strings.TrimSpace(identity.PrivateKey); privateKey != "" { |
| 407 | stored.PrivateKey = privateKey |
| 408 | } |
| 409 | if mnemonic := normalizeMnemonic(identity.Mnemonic); mnemonic != "" { |
| 410 | stored.Mnemonic = mnemonic |
| 411 | } |
| 412 | if derivationPath := strings.TrimSpace(identity.DerivationPath); derivationPath != "" { |
| 413 | stored.DerivationPath = derivationPath |
| 414 | } |
| 415 | if tokenSecret := strings.TrimSpace(identity.TokenSecret); tokenSecret != "" { |
| 416 | stored.TokenSecret = tokenSecret |
| 417 | } |
| 418 | if strings.TrimSpace(stored.PrivateKey) == "" { |
| 419 | return types.Identity{}, false, errors.New("stored identity private key is required") |
| 420 | } |
| 421 | if err := saveIdentity(path, stored); err != nil { |
| 422 | return types.Identity{}, false, fmt.Errorf("persist identity: %w", err) |
| 423 | } |
| 424 | loaded, err := loadIdentity(path) |
| 425 | if err != nil { |
| 426 | return types.Identity{}, false, fmt.Errorf("load identity: %w", err) |
| 427 | } |
| 428 | return loaded, false, nil |
| 429 | case !errors.Is(err, os.ErrNotExist): |
| 430 | return types.Identity{}, false, fmt.Errorf("load identity: %w", err) |
| 431 | } |
| 432 | |
| 433 | created := identity.Copy() |
| 434 | if strings.TrimSpace(created.Mnemonic) != "" || strings.TrimSpace(created.DerivationPath) != "" { |
| 435 | created, err = normalizeStoredIdentity(created) |
| 436 | if err != nil { |
| 437 | return types.Identity{}, false, fmt.Errorf("resolve identity mnemonic: %w", err) |
| 438 | } |
| 439 | } else { |
| 440 | generated, err := ResolveSecp256k1Identity(created.PrivateKey) |
| 441 | if err != nil { |
| 442 | return types.Identity{}, false, fmt.Errorf("generate identity: %w", err) |
| 443 | } |
| 444 | if strings.TrimSpace(created.Address) == "" { |
| 445 | created.Address = generated.Address |
| 446 | } |
| 447 | if strings.TrimSpace(created.PublicKey) == "" { |
| 448 | created.PublicKey = generated.PublicKey |
| 449 | } |
| 450 | created.PrivateKey = generated.PrivateKey |
| 451 | } |
| 452 | if strings.TrimSpace(created.TokenSecret) == "" { |
| 453 | created, err = ensureTokenSecret(created) |
| 454 | if err != nil { |
| 455 | return types.Identity{}, false, err |
| 456 | } |
| 457 | } |
| 458 | if err := saveIdentity(path, created); err != nil { |
| 459 | return types.Identity{}, false, fmt.Errorf("persist identity: %w", err) |
| 460 | } |
| 461 | loaded, err := loadIdentity(path) |
| 462 | if err != nil { |
| 463 | return types.Identity{}, false, fmt.Errorf("load identity: %w", err) |
| 464 | } |
| 465 | return loaded, true, nil |
| 466 | } |
| 467 | |
| 468 | func ResolveListenerIdentity(baseIdentity types.Identity, target, identityPath, identityJSON string) (types.Identity, bool, error) { |
| 469 | identityPath = strings.TrimSpace(identityPath) |
| 470 | identityJSON = strings.TrimSpace(identityJSON) |
| 471 | resolvedName, err := resolveExposeName(baseIdentity.Name, target, identityPath, identityJSON) |
| 472 | if err != nil { |
| 473 | return types.Identity{}, false, err |
| 474 | } |
| 475 | baseIdentity.Name = resolvedName |
| 476 | if identityJSON != "" { |
| 477 | provided, err := parseIdentityJSON(identityJSON) |
| 478 | if err != nil { |
| 479 | return types.Identity{}, false, err |
| 480 | } |
| 481 | provided.Name = baseIdentity.Name |
| 482 | if identityPath != "" { |
| 483 | if err := saveIdentity(identityPath, provided); err != nil { |
| 484 | return types.Identity{}, false, fmt.Errorf("persist identity: %w", err) |
| 485 | } |
| 486 | provided, err = loadIdentity(identityPath) |
| 487 | if err != nil { |
| 488 | return types.Identity{}, false, fmt.Errorf("load identity: %w", err) |
| 489 | } |
| 490 | } |
| 491 | resolved, err := resolveLeaseIdentity(provided) |
| 492 | return resolved, false, err |
| 493 | } |
| 494 | if identityPath == "" { |
| 495 | resolved, err := resolveLeaseIdentity(baseIdentity) |
| 496 | return resolved, false, err |
| 497 | } |
| 498 | |
| 499 | loaded, created, err := loadOrCreateIdentity(identityPath, baseIdentity) |
| 500 | if err != nil { |
| 501 | return types.Identity{}, false, err |
| 502 | } |
| 503 | resolved, err := resolveLeaseIdentity(loaded) |
| 504 | if err != nil { |
| 505 | return types.Identity{}, false, err |
| 506 | } |
| 507 | return resolved, created, nil |
| 508 | } |
| 509 | |
| 510 | func LoadOrCreateRelayIdentity(path, rootHost string, discoveryEnabled bool) (types.RelayIdentity, error) { |
| 511 | path = resolveRelayIdentityPath(path) |
| 512 | if path == "" { |
| 513 | return types.RelayIdentity{}, errors.New("identity path is required") |
| 514 | } |
| 515 | rootHost = strings.TrimSpace(rootHost) |
| 516 | if normalizedRootHost := utils.PortalRootHost(rootHost); normalizedRootHost != "" { |
| 517 | rootHost = normalizedRootHost |
| 518 | } else { |
| 519 | rootHost = utils.NormalizeHostname(rootHost) |
| 520 | } |
| 521 | |
| 522 | stored, err := loadRelayIdentity(path) |
| 523 | switch { |
| 524 | case err == nil: |
| 525 | if rootHost != "" { |
| 526 | stored.Name = rootHost |
| 527 | } |
| 528 | |
| 529 | if err := populateRelayIdentity(&stored, discoveryEnabled); err != nil { |
| 530 | return types.RelayIdentity{}, err |
| 531 | } |
| 532 | if err := saveRelayIdentity(path, stored); err != nil { |
| 533 | return types.RelayIdentity{}, fmt.Errorf("persist identity: %w", err) |
| 534 | } |
| 535 | loaded, err := loadRelayIdentity(path) |
| 536 | if err != nil { |
| 537 | return types.RelayIdentity{}, fmt.Errorf("load identity: %w", err) |
| 538 | } |
| 539 | return loaded, nil |
| 540 | case !errors.Is(err, os.ErrNotExist): |
| 541 | return types.RelayIdentity{}, fmt.Errorf("load identity: %w", err) |
| 542 | } |
| 543 | |
| 544 | created := types.RelayIdentity{ |
| 545 | Identity: types.Identity{Name: rootHost}, |
| 546 | } |
| 547 | generated, err := ResolveSecp256k1Identity(created.PrivateKey) |
| 548 | if err != nil { |
| 549 | return types.RelayIdentity{}, fmt.Errorf("generate identity: %w", err) |
| 550 | } |
| 551 | if strings.TrimSpace(created.Address) == "" { |
| 552 | created.Address = generated.Address |
| 553 | } |
| 554 | if strings.TrimSpace(created.PublicKey) == "" { |
| 555 | created.PublicKey = generated.PublicKey |
| 556 | } |
| 557 | created.PrivateKey = generated.PrivateKey |
| 558 | created.Identity, err = ensureTokenSecret(created.Identity) |
| 559 | if err != nil { |
| 560 | return types.RelayIdentity{}, err |
| 561 | } |
| 562 | |
| 563 | if err := populateRelayIdentity(&created, discoveryEnabled); err != nil { |
| 564 | return types.RelayIdentity{}, err |
| 565 | } |
| 566 | if err := saveRelayIdentity(path, created); err != nil { |
| 567 | return types.RelayIdentity{}, fmt.Errorf("persist identity: %w", err) |
| 568 | } |
| 569 | loaded, err := loadRelayIdentity(path) |
| 570 | if err != nil { |
| 571 | return types.RelayIdentity{}, fmt.Errorf("load identity: %w", err) |
| 572 | } |
| 573 | return loaded, nil |
| 574 | } |
| 575 | |
| 576 | func populateRelayIdentity(identity *types.RelayIdentity, discoveryEnabled bool) error { |
| 577 | if identity == nil { |
| 578 | return errors.New("relay identity is required") |
| 579 | } |
| 580 | baseIdentity, err := ensureTokenSecret(identity.Identity) |
| 581 | if err != nil { |
| 582 | return err |
| 583 | } |
| 584 | identity.Identity = baseIdentity |
| 585 | |
| 586 | if discoveryEnabled && strings.TrimSpace(identity.WireGuardPrivateKey) == "" { |
| 587 | var err error |
| 588 | wireGuardPrivateKey, err := GenerateWireGuardPrivateKey() |
| 589 | if err != nil { |
| 590 | return fmt.Errorf("generate relay wireguard private key: %w", err) |
| 591 | } |
| 592 | identity.WireGuardPrivateKey = wireGuardPrivateKey |
| 593 | } |
| 594 | |
| 595 | if strings.TrimSpace(identity.EncryptedClientHelloSeed) == "" { |
| 596 | identity.EncryptedClientHelloSeed = utils.RandomID("") |
| 597 | } |
| 598 | |
| 599 | return nil |
| 600 | } |
| 601 | |
| 602 | func normalizeIdentityKey(raw string) string { |
| 603 | key := strings.ToLower(strings.TrimSpace(raw)) |
| 604 | if key == "" { |
| 605 | return "" |
| 606 | } |
| 607 | name, address, ok := strings.Cut(key, types.IdentityKeySeparator) |
| 608 | if !ok || name == "" || address == "" { |
| 609 | return "" |
| 610 | } |
| 611 | return name + types.IdentityKeySeparator + address |
| 612 | } |
| 613 | |
| 614 | func NormalizeIdentityKeys(inputs []string) []string { |
| 615 | return normalizeUniqueStrings(inputs, normalizeIdentityKey) |
| 616 | } |
| 617 | |
| 618 | func NormalizeIdentityKeyBPS(inputs map[string]int64) map[string]int64 { |
| 619 | if len(inputs) == 0 { |
| 620 | return nil |
| 621 | } |
| 622 | |
| 623 | out := make(map[string]int64, len(inputs)) |
| 624 | for input, bps := range inputs { |
| 625 | key := normalizeIdentityKey(input) |
| 626 | if key == "" || bps <= 0 { |
| 627 | continue |
| 628 | } |
| 629 | out[key] = bps |
| 630 | } |
| 631 | if len(out) == 0 { |
| 632 | return nil |
| 633 | } |
| 634 | return out |
| 635 | } |
| 636 | |
| 637 | func resolveExposeName(name, target, identityPath, identityJSON string) (string, error) { |
| 638 | if name = strings.TrimSpace(name); name != "" { |
| 639 | return name, nil |
| 640 | } |
| 641 | if identityJSON = strings.TrimSpace(identityJSON); identityJSON != "" { |
| 642 | providedIdentity, err := parseIdentityJSON(identityJSON) |
| 643 | if err != nil { |
| 644 | return "", err |
| 645 | } |
| 646 | if name := strings.TrimSpace(providedIdentity.Name); name != "" { |
| 647 | return name, nil |
| 648 | } |
| 649 | } |
| 650 | if identityPath = strings.TrimSpace(identityPath); identityPath != "" { |
| 651 | storedIdentity, err := loadIdentity(identityPath) |
| 652 | switch { |
| 653 | case err == nil: |
| 654 | if name := strings.TrimSpace(storedIdentity.Name); name != "" { |
| 655 | return name, nil |
| 656 | } |
| 657 | case !errors.Is(err, os.ErrNotExist): |
| 658 | return "", err |
| 659 | } |
| 660 | } |
| 661 | |
| 662 | return utils.DefaultExposeName(target, utils.RandomID("cli_")) |
| 663 | } |
| 664 | |
| 665 | func resolveLeaseIdentity(identity types.Identity) (types.Identity, error) { |
| 666 | resolved, err := normalizeStoredIdentity(identity) |
| 667 | if err != nil { |
| 668 | return types.Identity{}, err |
| 669 | } |
| 670 | |
| 671 | name, err := utils.NormalizeDNSLabel(resolved.Name) |
| 672 | if err != nil { |
| 673 | return types.Identity{}, err |
| 674 | } |
| 675 | resolved.Name = name |
| 676 | |
| 677 | signingIdentity, err := ResolveSecp256k1Identity(resolved.PrivateKey) |
| 678 | if err != nil { |
| 679 | return types.Identity{}, err |
| 680 | } |
| 681 | if resolved.Address == "" { |
| 682 | resolved.Address = signingIdentity.Address |
| 683 | } else { |
| 684 | address, err := NormalizeEVMAddress(resolved.Address) |
| 685 | if err != nil { |
| 686 | return types.Identity{}, err |
| 687 | } |
| 688 | if address != signingIdentity.Address { |
| 689 | return types.Identity{}, errors.New("identity address does not match private key") |
| 690 | } |
| 691 | resolved.Address = address |
| 692 | } |
| 693 | |
| 694 | resolved.PublicKey = signingIdentity.PublicKey |
| 695 | resolved.PrivateKey = signingIdentity.PrivateKey |
| 696 | return ensureTokenSecret(resolved) |
| 697 | } |
| 698 | |
| 699 | func normalizeUniqueStrings(inputs []string, normalize func(string) string) []string { |
| 700 | if len(inputs) == 0 { |
| 701 | return nil |
| 702 | } |
| 703 | |
| 704 | out := make([]string, 0, len(inputs)) |
| 705 | seen := make(map[string]struct{}, len(inputs)) |
| 706 | for _, input := range inputs { |
| 707 | normalized := normalize(input) |
| 708 | if normalized == "" { |
| 709 | continue |
| 710 | } |
| 711 | if _, ok := seen[normalized]; ok { |
| 712 | continue |
| 713 | } |
| 714 | seen[normalized] = struct{}{} |
| 715 | out = append(out, normalized) |
| 716 | } |
| 717 | if len(out) == 0 { |
| 718 | return nil |
| 719 | } |
| 720 | return out |
| 721 | } |