package yuhticket import ( "crypto/hmac" "crypto/sha256" "encoding/base64" "encoding/json" "fmt" "os" "path/filepath" "strings" "sync" "time" ) const Prefix = "YHT1" // Claims 宇恒签发的短时凭票(仅用于智建 sync 绑定/免登录换票)。 type Claims struct { Iss string `json:"iss"` Aud string `json:"aud"` Phone string `json:"phone"` HostKey string `json:"host_key"` Name string `json:"name,omitempty"` LocalDatabaseID string `json:"local_database_id,omitempty"` Exp int64 `json:"exp"` JTI string `json:"jti"` Scope string `json:"scope,omitempty"` // sync_bind YuhengUserID string `json:"yuheng_user_id,omitempty"` } type VerifyOpts struct { Secret string Issuer string // default yuheng Audience string // default aijianzhan Now time.Time MaxSkew time.Duration // clock skew; default 30s } // Sign 供联调/测试;生产由宇恒侧用同一 Secret 签发。 func Sign(secret string, c Claims) (string, error) { secret = strings.TrimSpace(secret) if secret == "" { return "", fmt.Errorf("ticket secret empty") } if c.JTI == "" || c.Phone == "" || c.HostKey == "" || c.Exp == 0 { return "", fmt.Errorf("jti/phone/host_key/exp required") } if c.Iss == "" { c.Iss = "yuheng" } if c.Aud == "" { c.Aud = "aijianzhan" } if c.Scope == "" { c.Scope = "sync_bind" } raw, err := json.Marshal(c) if err != nil { return "", err } payload := base64.RawURLEncoding.EncodeToString(raw) mac := hmac.New(sha256.New, []byte(secret)) _, _ = mac.Write([]byte(Prefix + "." + payload)) sig := base64.RawURLEncoding.EncodeToString(mac.Sum(nil)) return Prefix + "." + payload + "." + sig, nil } func Verify(ticket string, opts VerifyOpts) (*Claims, error) { ticket = strings.TrimSpace(ticket) secret := strings.TrimSpace(opts.Secret) if secret == "" { return nil, fmt.Errorf("宇恒凭票未配置(Agent.YuhengTicket.Secret)") } parts := strings.Split(ticket, ".") if len(parts) != 3 || parts[0] != Prefix { return nil, fmt.Errorf("invalid ticket format") } mac := hmac.New(sha256.New, []byte(secret)) _, _ = mac.Write([]byte(parts[0] + "." + parts[1])) want := mac.Sum(nil) got, err := base64.RawURLEncoding.DecodeString(parts[2]) if err != nil || !hmac.Equal(want, got) { return nil, fmt.Errorf("ticket signature invalid") } raw, err := base64.RawURLEncoding.DecodeString(parts[1]) if err != nil { return nil, fmt.Errorf("ticket payload invalid") } var c Claims if err := json.Unmarshal(raw, &c); err != nil { return nil, fmt.Errorf("ticket claims invalid") } iss := strings.TrimSpace(opts.Issuer) if iss == "" { iss = "yuheng" } aud := strings.TrimSpace(opts.Audience) if aud == "" { aud = "aijianzhan" } if c.Iss != iss { return nil, fmt.Errorf("ticket issuer not allowed") } if c.Aud != aud { return nil, fmt.Errorf("ticket audience mismatch") } if c.Scope != "" && c.Scope != "sync_bind" { return nil, fmt.Errorf("ticket scope not allowed") } now := opts.Now if now.IsZero() { now = time.Now().UTC() } skew := opts.MaxSkew if skew <= 0 { skew = 30 * time.Second } exp := time.Unix(c.Exp, 0).UTC() if now.After(exp.Add(skew)) { return nil, fmt.Errorf("ticket expired") } // 拒绝过远未来的 exp(防永久票) if exp.After(now.Add(10 * time.Minute)) { return nil, fmt.Errorf("ticket exp too far") } if strings.TrimSpace(c.Phone) == "" || strings.TrimSpace(c.HostKey) == "" || strings.TrimSpace(c.JTI) == "" { return nil, fmt.Errorf("ticket missing phone/host_key/jti") } return &c, nil } // JTIStore 防重放(短 TTL)。 type JTIStore struct { mu sync.Mutex path string seen map[string]int64 // jti -> exp unix } func NewJTIStore(dir string) (*JTIStore, error) { if err := os.MkdirAll(dir, 0o755); err != nil { return nil, err } s := &JTIStore{path: filepath.Join(dir, "yuheng_ticket_jti.json"), seen: map[string]int64{}} _ = s.load() return s, nil } func (s *JTIStore) load() error { b, err := os.ReadFile(s.path) if err != nil { if os.IsNotExist(err) { return nil } return err } return json.Unmarshal(b, &s.seen) } func (s *JTIStore) save() error { b, err := json.Marshal(s.seen) if err != nil { return err } tmp := s.path + ".tmp" if err := os.WriteFile(tmp, b, 0o644); err != nil { return err } return os.Rename(tmp, s.path) } // Consume 若 jti 已用则报错;否则记入至 exp。 func (s *JTIStore) Consume(jti string, expUnix int64) error { s.mu.Lock() defer s.mu.Unlock() now := time.Now().Unix() for k, exp := range s.seen { if exp < now { delete(s.seen, k) } } if _, ok := s.seen[jti]; ok { return fmt.Errorf("ticket already used") } s.seen[jti] = expUnix return s.save() }