Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
41 changes: 32 additions & 9 deletions internal/store/coverage_contracts_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -1369,19 +1369,42 @@ func TestMarkdownImportReportsTransactionalWriteFailures(t *testing.T) {
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
s := newStore(t)
path := filepath.Join(t.TempDir(), "u.md")
if err := os.WriteFile(path, []byte(tt.markdown), 0o600); err != nil {
t.Fatal(err)
}
mustExecCoverage(t, s, tt.trigger)
if _, err := s.importBoard(path, "u"); err == nil {
t.Fatal("importBoard returned nil error")
}
testMarkdownImportRollback(t, tt.trigger, tt.markdown)
})
}
}

func testMarkdownImportRollback(t *testing.T, trigger, markdown string) {
t.Helper()
s := newStore(t)
path := filepath.Join(t.TempDir(), "u.md")
if err := os.WriteFile(path, []byte(markdown), 0o600); err != nil {
t.Fatal(err)
}
mustExecCoverage(t, s, trigger)
if _, err := s.importBoard(path, "u"); err == nil {
t.Fatal("importBoard returned nil error")
}
requireEmptyMarkdownImportState(t, s)
}

func requireEmptyMarkdownImportState(t *testing.T, s *Store) {
t.Helper()
var tasks, labels, metadata int
if err := s.db.QueryRow(`SELECT COUNT(*) FROM tasks WHERE user = 'u'`).Scan(&tasks); err != nil {
t.Fatal(err)
}
if err := s.db.QueryRow(`SELECT COUNT(*) FROM labels WHERE user = 'u'`).Scan(&labels); err != nil {
t.Fatal(err)
}
if err := s.db.QueryRow(`SELECT COUNT(*) FROM meta WHERE k IN ('board_title:u', 'imported:u')`).Scan(&metadata); err != nil {
t.Fatal(err)
}
if tasks != 0 || labels != 0 || metadata != 0 {
t.Fatalf("failed import left tasks=%d labels=%d metadata=%d", tasks, labels, metadata)
}
}

func TestSimilaritySearchRejectsUnscannableFTSRows(t *testing.T) {
t.Run("task result", func(t *testing.T) {
s := newStore(t)
Expand Down
164 changes: 106 additions & 58 deletions internal/store/forge.go
Original file line number Diff line number Diff line change
Expand Up @@ -22,6 +22,20 @@ type ForgeSource struct {
CreatedAt time.Time
}

type forgeSourceInput struct {
name string
kind string
baseURL *string
pat *string
}

type forgeSourceState struct {
baseURL string
patEnc []byte
createdAt string
exists bool
}

// ForgeSources returns the scope's configured forges ordered by canonical name.
// It checks only whether ciphertext exists; listing never decrypts a PAT.
func (s *Store) ForgeSources(scope string) ([]ForgeSource, error) {
Expand Down Expand Up @@ -78,69 +92,19 @@ ORDER BY name`, scope)
// origin changes without a PAT in the same call, the old ciphertext is cleared
// atomically and tokenCleared reports that the caller must request a new PAT.
func (s *Store) SetForgeSource(scope, name, kind string, baseURL, pat *string) (tokenCleared bool, err error) {
name, err = normalizeForgeSourceName(name)
input, err := normalizeForgeSourceInput(name, kind, baseURL, pat)
if err != nil {
return false, err
}
if kind != "gitlab" && kind != "github" {
return false, errors.New("store: invalid forge kind")
}

var normalizedBaseURL *string
if baseURL != nil {
base, err := normalizeForgeBaseURL(*baseURL)
if err != nil {
return false, err
}
normalizedBaseURL = &base
}

err = s.withTx(func(tx *sql.Tx) error {
var base, createdAt string
var enc []byte
exists := true
switch err := tx.QueryRow(`
SELECT base_url, pat_enc, created_at
FROM forge_sources
WHERE scope = ? AND name = ?`, scope, name).Scan(&base, &enc, &createdAt); {
case errors.Is(err, sql.ErrNoRows):
exists = false
case err != nil:
return fmt.Errorf("store: read forge source: %w", err)
}
if exists {
clean := stripURLSuffix(base)
if clean != base {
if _, err := tx.Exec(`UPDATE forge_sources SET base_url = ? WHERE scope = ? AND name = ?`, clean, scope, name); err != nil {
return fmt.Errorf("store: repair forge source URL: %w", err)
}
base = clean
}
}

if !exists {
if normalizedBaseURL == nil {
return errors.New("store: forge base URL is required")
}
createdAt = time.Now().UTC().Format(time.RFC3339Nano)
}
if normalizedBaseURL != nil {
if pat == nil && len(enc) > 0 && !SameAIOrigin(base, *normalizedBaseURL) {
enc = nil
tokenCleared = true
}
base = *normalizedBaseURL
state, err := loadForgeSourceState(tx, scope, input.name)
if err != nil {
return err
}
if pat != nil {
if *pat == "" {
enc = nil
} else {
sealed, err := s.seal([]byte(*pat))
if err != nil {
return err
}
enc = sealed
}
tokenCleared, err = s.patchForgeSourceState(&state, input)
if err != nil {
return err
}

if _, err := tx.Exec(`
Expand All @@ -150,7 +114,7 @@ ON CONFLICT(scope, name) DO UPDATE SET
kind = excluded.kind,
base_url = excluded.base_url,
pat_enc = excluded.pat_enc`,
scope, name, kind, base, enc, createdAt); err != nil {
scope, input.name, input.kind, state.baseURL, state.patEnc, state.createdAt); err != nil {
return fmt.Errorf("store: write forge source: %w", err)
}
return nil
Expand All @@ -161,6 +125,90 @@ ON CONFLICT(scope, name) DO UPDATE SET
return tokenCleared, nil
}

func normalizeForgeSourceInput(name, kind string, baseURL, pat *string) (forgeSourceInput, error) {
name, err := normalizeForgeSourceName(name)
if err != nil {
return forgeSourceInput{}, err
}
if kind != "gitlab" && kind != "github" {
return forgeSourceInput{}, errors.New("store: invalid forge kind")
}
if baseURL == nil {
return forgeSourceInput{name: name, kind: kind, pat: pat}, nil
}
base, err := normalizeForgeBaseURL(*baseURL)
if err != nil {
return forgeSourceInput{}, err
}
return forgeSourceInput{name: name, kind: kind, baseURL: &base, pat: pat}, nil
}

func loadForgeSourceState(tx *sql.Tx, scope, name string) (forgeSourceState, error) {
state := forgeSourceState{exists: true}
err := tx.QueryRow(`
SELECT base_url, pat_enc, created_at
FROM forge_sources
WHERE scope = ? AND name = ?`, scope, name).Scan(&state.baseURL, &state.patEnc, &state.createdAt)
if errors.Is(err, sql.ErrNoRows) {
state.exists = false
return state, nil
}
if err != nil {
return forgeSourceState{}, fmt.Errorf("store: read forge source: %w", err)
}
clean := stripURLSuffix(state.baseURL)
if clean == state.baseURL {
return state, nil
}
if _, err := tx.Exec(`UPDATE forge_sources SET base_url = ? WHERE scope = ? AND name = ?`, clean, scope, name); err != nil {
return forgeSourceState{}, fmt.Errorf("store: repair forge source URL: %w", err)
}
state.baseURL = clean
return state, nil
}

func (s *Store) patchForgeSourceState(state *forgeSourceState, input forgeSourceInput) (bool, error) {
if !state.exists {
if input.baseURL == nil {
return false, errors.New("store: forge base URL is required")
}
state.createdAt = time.Now().UTC().Format(time.RFC3339Nano)
}
tokenCleared := patchForgeBaseURL(state, input.baseURL, input.pat)
if err := s.patchForgePAT(state, input.pat); err != nil {
return false, err
}
return tokenCleared, nil
}

func patchForgeBaseURL(state *forgeSourceState, baseURL, pat *string) bool {
if baseURL == nil {
return false
}
tokenCleared := pat == nil && len(state.patEnc) > 0 && !SameAIOrigin(state.baseURL, *baseURL)
if tokenCleared {
state.patEnc = nil
}
state.baseURL = *baseURL
return tokenCleared
}

func (s *Store) patchForgePAT(state *forgeSourceState, pat *string) error {
if pat == nil {
return nil
}
if *pat == "" {
state.patEnc = nil
return nil
}
sealed, err := s.seal([]byte(*pat))
if err != nil {
return err
}
state.patEnc = sealed
return nil
}

// DeleteForgeSource deletes one scoped source. Deleting a missing source is a
// successful no-op.
func (s *Store) DeleteForgeSource(scope, name string) error {
Expand Down
79 changes: 47 additions & 32 deletions internal/store/migrate.go
Original file line number Diff line number Diff line change
Expand Up @@ -57,47 +57,62 @@ func (s *Store) ImportMarkdownDir(dir string) (int, error) {
func (s *Store) importBoard(path, user string) (bool, error) {
did := false
err := s.withTx(func(tx *sql.Tx) error {
var v string
switch err := tx.QueryRow(`SELECT v FROM meta WHERE k = ?`, importKey(user)).Scan(&v); {
case err == nil:
return nil // already handled once; never reimport
case !errors.Is(err, sql.ErrNoRows):
return fmt.Errorf("store: import flag: %w", err)
}
var n int
if err := tx.QueryRow(`SELECT COUNT(*) FROM tasks WHERE user = ?`, user).Scan(&n); err != nil {
return fmt.Errorf("store: count tasks: %w", err)
}
if n > 0 {
return setMeta(tx, importKey(user), "skipped")
eligible, err := importBoardEligible(tx, user)
if err != nil || !eligible {
return err
}
data, err := os.ReadFile(path)
if err != nil {
return fmt.Errorf("store: read %s: %w", path, err)
}
b := board.Parse(string(data))
now := time.Now().UTC()
pos := map[board.Status]int{}
for _, t := range b.Tasks {
t.ID = uuid.NewString()
t.CreatedAt, t.MovedAt = now, now
t.Position = pos[t.Status]
pos[t.Status]++
if err := insertTask(tx, user, t); err != nil {
return err
}
if err := s.upsertLabels(tx, user, t.Tags); err != nil {
return err
}
}
if err := setMeta(tx, titleKey(user), b.Title); err != nil {
return err
}
if err := setMeta(tx, importKey(user), "imported"); err != nil {
if err := s.insertImportedBoard(tx, user, board.Parse(string(data))); err != nil {
return err
}
did = true
return nil
})
return did, err
}

func importBoardEligible(tx *sql.Tx, user string) (bool, error) {
var value string
err := tx.QueryRow(`SELECT v FROM meta WHERE k = ?`, importKey(user)).Scan(&value)
if err == nil {
return false, nil
}
if !errors.Is(err, sql.ErrNoRows) {
return false, fmt.Errorf("store: import flag: %w", err)
}
var taskCount int
if err := tx.QueryRow(`SELECT COUNT(*) FROM tasks WHERE user = ?`, user).Scan(&taskCount); err != nil {
return false, fmt.Errorf("store: count tasks: %w", err)
}
if taskCount == 0 {
return true, nil
}
if err := setMeta(tx, importKey(user), "skipped"); err != nil {
return false, err
}
return false, nil
}

func (s *Store) insertImportedBoard(tx *sql.Tx, user string, imported board.Board) error {
now := time.Now().UTC()
positions := map[board.Status]int{}
for _, task := range imported.Tasks {
task.ID = uuid.NewString()
task.CreatedAt, task.MovedAt = now, now
task.Position = positions[task.Status]
positions[task.Status]++
if err := insertTask(tx, user, task); err != nil {
return err
}
if err := s.upsertLabels(tx, user, task.Tags); err != nil {
return err
}
}
if err := setMeta(tx, titleKey(user), imported.Title); err != nil {
return err
}
return setMeta(tx, importKey(user), "imported")
}
Loading