| @@ -5,6 +5,7 @@ import ( |
| 5 | "errors" |
5 | "errors" |
| 6 | "fmt" |
6 | "fmt" |
| 7 | "strings" |
7 | "strings" |
| |
8 | "time" |
| 8 | ) |
9 | ) |
| 9 | |
10 | |
| 10 | type User struct { |
11 | type User struct { |
| @@ -24,8 +25,14 @@ type SSHKey struct { |
| 24 | Scope string |
25 | Scope string |
| 25 | Label string // "" when the key was added with no name |
26 | Label string // "" when the key was added with no name |
| 26 | CreatedAt string |
27 | CreatedAt string |
| 27 | LastUsedAt string // "" when the key has never authenticated |
28 | LastUsedAt string // "" when the key has never authenticated |
| 28 | CreatedBy string // name of the API token that added the key; "" for none. ListSSHKeys only. |
29 | CreatedBy string // name of the API token that added the key; "" for none. ListSSHKeys only. |
| |
30 | ExpiresAt *time.Time // nil when the key never expires |
| |
31 | } |
| |
32 | |
| |
33 | // Expired reports whether the key has lapsed at now. |
| |
34 | func (k SSHKey) Expired(now time.Time) bool { |
| |
35 | return k.ExpiresAt != nil && !k.ExpiresAt.After(now) |
| 29 | } |
36 | } |
| 30 | |
37 | |
| 31 | // ErrDuplicateKey carries the exact user-facing message from the spec. It |
38 | // ErrDuplicateKey carries the exact user-facing message from the spec. It |
| @@ -273,7 +280,8 @@ func (s *Store) UserByID(id int64) (User, error) { |
| 273 | |
280 | |
| 274 | // KeyOrigin is how a key came to be. |
281 | // KeyOrigin is how a key came to be. |
| 275 | type KeyOrigin struct { |
282 | type KeyOrigin struct { |
| 276 | CreatedByToken int64 // the API token that added it; 0 for none |
283 | CreatedByToken int64 // the API token that added it; 0 for none |
| |
284 | ExpiresAt *time.Time // when it stops authenticating; nil for never |
| 277 | } |
285 | } |
| 278 | |
286 | |
| 279 | // AddSSHKey registers a key and bumps the key epoch in one transaction. |
287 | // AddSSHKey registers a key and bumps the key epoch in one transaction. |
| @@ -288,9 +296,13 @@ func (s *Store) AddSSHKeyFrom(userID int64, fingerprint, algo string, blob []byt |
| 288 | return err |
296 | return err |
| 289 | } |
297 | } |
| 290 | defer tx.Rollback() |
298 | defer tx.Rollback() |
| |
299 | var exp any |
| |
300 | if o.ExpiresAt != nil { |
| |
301 | exp = fmtTime(*o.ExpiresAt) |
| |
302 | } |
| 291 | if _, err := tx.Exec( |
303 | if _, err := tx.Exec( |
| 292 | "INSERT INTO ssh_keys (user_id, fingerprint, algo, blob, scope, label, created_by_token) VALUES (?, ?, ?, ?, ?, ?, ?)", |
304 | "INSERT INTO ssh_keys (user_id, fingerprint, algo, blob, scope, label, created_by_token, expires_at) VALUES (?, ?, ?, ?, ?, ?, ?, ?)", |
| 293 | userID, fingerprint, algo, blob, scope, label, nullID(o.CreatedByToken)); err != nil { |
305 | userID, fingerprint, algo, blob, scope, label, nullID(o.CreatedByToken), exp); err != nil { |
| 294 | if isUniqueErr(err) { |
306 | if isUniqueErr(err) { |
| 295 | return ErrDuplicateKey |
307 | return ErrDuplicateKey |
| 296 | } |
308 | } |
| @@ -343,19 +355,21 @@ func (s *Store) SetSSHKeyLabel(userID int64, fingerprint, label string) error { |
| 343 | |
355 | |
| 344 | func (s *Store) SSHKeyByFingerprint(fingerprint string) (SSHKey, error) { |
356 | func (s *Store) SSHKeyByFingerprint(fingerprint string) (SSHKey, error) { |
| 345 | var k SSHKey |
357 | var k SSHKey |
| |
358 | var exp sql.NullString |
| 346 | err := s.DB.QueryRow( |
359 | err := s.DB.QueryRow( |
| 347 | "SELECT id, user_id, fingerprint, algo, blob, scope, label FROM ssh_keys WHERE fingerprint = ?", |
360 | "SELECT id, user_id, fingerprint, algo, blob, scope, label, expires_at FROM ssh_keys WHERE fingerprint = ?", |
| 348 | fingerprint).Scan(&k.ID, &k.UserID, &k.Fingerprint, &k.Algo, &k.Blob, &k.Scope, &k.Label) |
361 | fingerprint).Scan(&k.ID, &k.UserID, &k.Fingerprint, &k.Algo, &k.Blob, &k.Scope, &k.Label, &exp) |
| 349 | if errors.Is(err, sql.ErrNoRows) { |
362 | if errors.Is(err, sql.ErrNoRows) { |
| 350 | return k, ErrNotFound |
363 | return k, ErrNotFound |
| 351 | } |
364 | } |
| |
365 | k.ExpiresAt = parseTime(exp) |
| 352 | return k, err |
366 | return k, err |
| 353 | } |
367 | } |
| 354 | |
368 | |
| 355 | func (s *Store) ListSSHKeys(userID int64) ([]SSHKey, error) { |
369 | func (s *Store) ListSSHKeys(userID int64) ([]SSHKey, error) { |
| 356 | rows, err := s.DB.Query( |
370 | rows, err := s.DB.Query( |
| 357 | `SELECT k.id, k.user_id, k.fingerprint, k.algo, k.blob, k.scope, k.label, k.created_at, |
371 | `SELECT k.id, k.user_id, k.fingerprint, k.algo, k.blob, k.scope, k.label, k.created_at, |
| 358 | COALESCE(k.last_used_at, ''), COALESCE(t.name, '') |
372 | COALESCE(k.last_used_at, ''), COALESCE(t.name, ''), k.expires_at |
| 359 | FROM ssh_keys k LEFT JOIN api_tokens t ON t.id = k.created_by_token |
373 | FROM ssh_keys k LEFT JOIN api_tokens t ON t.id = k.created_by_token |
| 360 | WHERE k.user_id = ? ORDER BY k.id`, |
374 | WHERE k.user_id = ? ORDER BY k.id`, |
| 361 | userID) |
375 | userID) |
| @@ -366,9 +380,11 @@ func (s *Store) ListSSHKeys(userID int64) ([]SSHKey, error) { |
| 366 | var keys []SSHKey |
380 | var keys []SSHKey |
| 367 | for rows.Next() { |
381 | for rows.Next() { |
| 368 | var k SSHKey |
382 | var k SSHKey |
| 369 | if err := rows.Scan(&k.ID, &k.UserID, &k.Fingerprint, &k.Algo, &k.Blob, &k.Scope, &k.Label, &k.CreatedAt, &k.LastUsedAt, &k.CreatedBy); err != nil { |
383 | var exp sql.NullString |
| |
384 | if err := rows.Scan(&k.ID, &k.UserID, &k.Fingerprint, &k.Algo, &k.Blob, &k.Scope, &k.Label, &k.CreatedAt, &k.LastUsedAt, &k.CreatedBy, &exp); err != nil { |
| 370 | return nil, err |
385 | return nil, err |
| 371 | } |
386 | } |
| |
387 | k.ExpiresAt = parseTime(exp) |
| 372 | keys = append(keys, k) |
388 | keys = append(keys, k) |
| 373 | } |
389 | } |
| 374 | return keys, rows.Err() |
390 | return keys, rows.Err() |
| @@ -523,19 +539,22 @@ func isUniqueErr(err error) bool { |
| 523 | |
539 | |
| 524 | func (s *Store) SSHKeyByID(id int64) (SSHKey, error) { |
540 | func (s *Store) SSHKeyByID(id int64) (SSHKey, error) { |
| 525 | var k SSHKey |
541 | var k SSHKey |
| |
542 | var exp sql.NullString |
| 526 | err := s.DB.QueryRow( |
543 | err := s.DB.QueryRow( |
| 527 | "SELECT id, user_id, fingerprint, algo, blob, scope, label FROM ssh_keys WHERE id = ?", |
544 | "SELECT id, user_id, fingerprint, algo, blob, scope, label, expires_at FROM ssh_keys WHERE id = ?", |
| 528 | id).Scan(&k.ID, &k.UserID, &k.Fingerprint, &k.Algo, &k.Blob, &k.Scope, &k.Label) |
545 | id).Scan(&k.ID, &k.UserID, &k.Fingerprint, &k.Algo, &k.Blob, &k.Scope, &k.Label, &exp) |
| 529 | if errors.Is(err, sql.ErrNoRows) { |
546 | if errors.Is(err, sql.ErrNoRows) { |
| 530 | return k, ErrNotFound |
547 | return k, ErrNotFound |
| 531 | } |
548 | } |
| |
549 | k.ExpiresAt = parseTime(exp) |
| 532 | return k, err |
550 | return k, err |
| 533 | } |
551 | } |
| 534 | |
552 | |
| 535 | // ListDeployKeys returns the deploy keys bound to a repository. |
553 | // ListDeployKeys returns the deploy keys bound to a repository. |
| 536 | func (s *Store) ListDeployKeys(repoID int64) ([]SSHKey, error) { |
554 | func (s *Store) ListDeployKeys(repoID int64) ([]SSHKey, error) { |
| 537 | rows, err := s.DB.Query( |
555 | rows, err := s.DB.Query( |
| 538 | "SELECT id, user_id, fingerprint, algo, blob, scope, label FROM ssh_keys WHERE scope LIKE 'deploy:' || ? || ':%' ORDER BY id", |
556 | `SELECT id, user_id, fingerprint, algo, blob, scope, label, COALESCE(last_used_at, ''), expires_at |
| |
557 | FROM ssh_keys WHERE scope LIKE 'deploy:' || ? || ':%' ORDER BY id`, |
| 539 | repoID) |
558 | repoID) |
| 540 | if err != nil { |
559 | if err != nil { |
| 541 | return nil, err |
560 | return nil, err |
| @@ -544,9 +563,11 @@ func (s *Store) ListDeployKeys(repoID int64) ([]SSHKey, error) { |
| 544 | var keys []SSHKey |
563 | var keys []SSHKey |
| 545 | for rows.Next() { |
564 | for rows.Next() { |
| 546 | var k SSHKey |
565 | var k SSHKey |
| 547 | if err := rows.Scan(&k.ID, &k.UserID, &k.Fingerprint, &k.Algo, &k.Blob, &k.Scope, &k.Label); err != nil { |
566 | var exp sql.NullString |
| |
567 | if err := rows.Scan(&k.ID, &k.UserID, &k.Fingerprint, &k.Algo, &k.Blob, &k.Scope, &k.Label, &k.LastUsedAt, &exp); err != nil { |
| 548 | return nil, err |
568 | return nil, err |
| 549 | } |
569 | } |
| |
570 | k.ExpiresAt = parseTime(exp) |
| 550 | keys = append(keys, k) |
571 | keys = append(keys, k) |
| 551 | } |
572 | } |
| 552 | return keys, rows.Err() |
573 | return keys, rows.Err() |