pico
created pr with
97.1
added 97.2
1: fdc255e = 1: fdc255e refactor: remove unused db methods
2: 3f855f6 = 2: 3f855f6 chore: add tests for postgres db impl
3: 7cd29a6 = 3: 7cd29a6 refactor: use sqlx interface
4: acd449e = 4: acd449e chore: add db tags
5: f6ed0d7 = 5: f6ed0d7 refactor: use sqlx
6: 29035c6 = 6: 29035c6 refactor: replace custom sql functions with sqlx
7: d4ea6d9 = 7: d4ea6d9 refactor: inline all sql queries
-: ------- > 8: 4bd37ed refactor: use `select * from` where possible
cmds
checkout latest patchset:
ssh pr.pico.sh print 97 | git am -3checkout any patchset in a patch request:
ssh pr.pico.sh print 97.[rev] | git am -3add changes to patch request:
git format-patch main --stdout | ssh pr.pico.sh pr add 97
Patchset
97.1
refactor: remove unused db methods
Eric Bower
chore: add tests for postgres db impl
2025-12-18T03:03:10ZEric Bower
refactor: use sqlx interface
2025-12-18T03:17:57ZEric Bower
chore: add db tags
2025-12-18T03:27:57ZEric Bower
→ refactor: use sqlx
2025-12-18T03:31:25ZEric Bower
refactor: replace custom sql functions with sqlx
2025-12-18T03:43:01ZEric Bower
refactor: inline all sql queries
2025-12-18T03:57:00ZEric Bower
2025-12-18T04:02:50Z
refactor: use sqlx
Eric Bower
2025-12-18T03:43:01ZSemantic diff summary
1 added,
31 modified,
6 signature changed,
0 removed
across 6 analyzed files
pkg/db/postgres/storage.go
-
chunklines 1-6modified -
chunklines 124-133modified -
function_declarationRegisterUsermodified -
method_declarationinsertPublicKeyWithTxadded -
method_declarationInsertPublicKeysignature changed -
method_declarationfindPublicKeymodified -
method_declarationFindKeysForUsermodified -
method_declarationFindUsermodified -
method_declarationFindUserByNamemodified -
method_declarationFindUserForTokenmodified -
method_declarationFindUsersmodified -
method_declarationremoveTagsForPostsignature changed -
method_declarationinsertTagsForPostsignature changed -
method_declarationReplaceTagsForPostmodified -
method_declarationinsertAliasesForPostsignature changed -
method_declarationremoveAliasesForPostsignature changed -
method_declarationReplaceAliasesForPostmodified -
method_declarationFindFeaturemodified -
function_declarationFindFeaturesForUsermodified -
function_declarationHasFeatureForUsermodified -
method_declarationFindFeedItemsByPostIDmodified -
method_declarationFindProjectByNamemodified -
method_declarationFindTokensForUsermodified -
function_declarationAddPicoPlusUsermodified -
method_declarationFindTunsEventLogsByAddrmodified -
method_declarationFindTunsEventLogsmodified -
method_declarationFindAccessLogsmodified -
method_declarationFindAccessLogsByPubkeymodified -
method_declarationFindPubkeysInAccessLogsmodified
+1
-1
pkg/apps/pico/file_handler.go
#
| ... | ... | @@ -247,7 +247,7 @@ func (h *UploadHandler) ProcessAuthorizedKeys(text []byte, logger *slog.Logger, | |
| 247 | 247 | _, _ = fmt.Fprintf(s.Stderr(), "adding pubkey (%s)\n", key) | |
| 248 | 248 | logger.Info("adding pubkey", "pubkey", key) | |
| 249 | 249 | ||
| 250 | - | err = dbpool.InsertPublicKey(user.ID, key, pk.Comment, nil) | |
| 250 | + | err = dbpool.InsertPublicKey(user.ID, key, pk.Comment) | |
| 251 | 251 | if err != nil { | |
| 252 | 252 | _, _ = fmt.Fprintf(s.Stderr(), "error: could not insert pubkey: %s (%s)\n", err.Error(), key) | |
| 253 | 253 | logger.Error("could not insert pubkey", "err", err.Error()) |
+1
-1
pkg/db/db.go
#
| ... | ... | @@ -390,7 +390,7 @@ var DenyList = []string{ | |
| 390 | 390 | type DB interface { | |
| 391 | 391 | RegisterUser(name, pubkey, comment string) (*User, error) | |
| 392 | 392 | UpdatePublicKey(pubkeyID, name string) (*PublicKey, error) | |
| 393 | - | InsertPublicKey(userID, pubkey, name string, tx *sql.Tx) error | |
| 393 | + | InsertPublicKey(userID, pubkey, name string) error | |
| 394 | 394 | FindKeysForUser(user *User) ([]*PublicKey, error) | |
| 395 | 395 | RemoveKeys(pubkeyIDs []string) error | |
| 396 | 396 |
+61
-299
pkg/db/postgres/storage.go
#
| ... | ... | @@ -125,10 +124,10 @@ var ( | |
| 125 | 124 | const ( | |
| 126 | 125 | sqlSelectPublicKey = `SELECT id, user_id, name, public_key, created_at FROM public_keys WHERE public_key = $1` | |
| 127 | 126 | sqlSelectPublicKeys = `SELECT id, user_id, name, public_key, created_at FROM public_keys WHERE user_id = $1 ORDER BY created_at ASC` | |
| 128 | - | sqlSelectUser = `SELECT id, name, created_at FROM app_users WHERE id = $1` | |
| 127 | + | sqlSelectUser = `SELECT id, COALESCE(name, '') as name, created_at FROM app_users WHERE id = $1` | |
| 129 | 128 | sqlSelectUserForName = `SELECT id, name, created_at FROM app_users WHERE name = $1` | |
| 130 | 129 | sqlSelectUserForNameAndKey = `SELECT app_users.id, app_users.name, app_users.created_at, public_keys.id as pk_id, public_keys.public_key, public_keys.created_at as pk_created_at FROM app_users LEFT JOIN public_keys ON public_keys.user_id = app_users.id WHERE app_users.name = $1 AND public_keys.public_key = $2` | |
| 131 | - | sqlSelectUsers = `SELECT id, name, created_at FROM app_users ORDER BY name ASC` | |
| 130 | + | sqlSelectUsers = `SELECT id, COALESCE(name, '') as name, created_at FROM app_users ORDER BY name ASC` | |
| 132 | 131 | ||
| 133 | 132 | sqlSelectUserForToken = ` | |
| 134 | 133 | SELECT app_users.id, app_users.name, app_users.created_at |
| ... | ... | @@ -313,30 +312,21 @@ func (me *PsqlDB) RegisterUser(username, pubkey, comment string) (*db.User, erro | |
| 313 | 312 | return nil, err | |
| 314 | 313 | } | |
| 315 | 314 | ||
| 316 | - | ctx := context.Background() | |
| 317 | - | tx, err := me.Db.BeginTx(ctx, nil) | |
| 315 | + | tx, err := me.Db.Beginx() | |
| 318 | 316 | if err != nil { | |
| 319 | 317 | return nil, err | |
| 320 | 318 | } | |
| 321 | 319 | defer func() { | |
| 322 | - | err = tx.Rollback() | |
| 323 | - | }() | |
| 324 | - | ||
| 325 | - | stmt, err := tx.Prepare(sqlInsertUser) | |
| 326 | - | if err != nil { | |
| 327 | - | return nil, err | |
| 328 | - | } | |
| 329 | - | defer func() { | |
| 330 | - | _ = stmt.Close() | |
| 320 | + | _ = tx.Rollback() | |
| 331 | 321 | }() | |
| 332 | 322 | ||
| 333 | 323 | var id string | |
| 334 | - | err = stmt.QueryRow(lowerName).Scan(&id) | |
| 324 | + | err = tx.QueryRow(sqlInsertUser, lowerName).Scan(&id) | |
| 335 | 325 | if err != nil { | |
| 336 | 326 | return nil, err | |
| 337 | 327 | } | |
| 338 | 328 | ||
| 339 | - | err = me.InsertPublicKey(id, pubkey, comment, tx) | |
| 329 | + | err = me.insertPublicKeyWithTx(id, pubkey, comment, tx) | |
| 340 | 330 | if err != nil { | |
| 341 | 331 | return nil, err | |
| 342 | 332 | } |
| ... | ... | @@ -349,23 +339,24 @@ func (me *PsqlDB) RegisterUser(username, pubkey, comment string) (*db.User, erro | |
| 349 | 339 | return me.FindUserForKey(username, pubkey) | |
| 350 | 340 | } | |
| 351 | 341 | ||
| 352 | - | func (me *PsqlDB) InsertPublicKey(userID, key, name string, tx *sql.Tx) error { | |
| 342 | + | func (me *PsqlDB) insertPublicKeyWithTx(userID, key, name string, tx *sqlx.Tx) error { | |
| 353 | 343 | pk, _ := me.findPublicKeyForKey(key) | |
| 354 | 344 | if pk != nil { | |
| 355 | 345 | return db.ErrPublicKeyTaken | |
| 356 | 346 | } | |
| 357 | 347 | query := `INSERT INTO public_keys (user_id, public_key, name) VALUES ($1, $2, $3)` | |
| 358 | - | var err error | |
| 359 | - | if tx != nil { | |
| 360 | - | _, err = tx.Exec(query, userID, key, name) | |
| 361 | - | } else { | |
| 362 | - | _, err = me.Db.Exec(query, userID, key, name) | |
| 363 | - | } | |
| 364 | - | if err != nil { | |
| 365 | - | return err | |
| 366 | - | } | |
| 348 | + | _, err := tx.Exec(query, userID, key, name) | |
| 349 | + | return err | |
| 350 | + | } | |
| 367 | 351 | ||
| 368 | - | return nil | |
| 352 | + | func (me *PsqlDB) InsertPublicKey(userID, key, name string) error { | |
| 353 | + | pk, _ := me.findPublicKeyForKey(key) | |
| 354 | + | if pk != nil { | |
| 355 | + | return db.ErrPublicKeyTaken | |
| 356 | + | } | |
| 357 | + | query := `INSERT INTO public_keys (user_id, public_key, name) VALUES ($1, $2, $3)` | |
| 358 | + | _, err := me.Db.Exec(query, userID, key, name) | |
| 359 | + | return err | |
| 369 | 360 | } | |
| 370 | 361 | ||
| 371 | 362 | func (me *PsqlDB) UpdatePublicKey(pubkeyID, name string) (*db.PublicKey, error) { |
| ... | ... | @@ -424,50 +415,19 @@ func (me *PsqlDB) findPublicKeyForKey(key string) (*db.PublicKey, error) { | |
| 424 | 415 | } | |
| 425 | 416 | ||
| 426 | 417 | func (me *PsqlDB) findPublicKey(pubkeyID string) (*db.PublicKey, error) { | |
| 427 | - | var keys []*db.PublicKey | |
| 428 | - | rs, err := me.Db.Query(`SELECT id, user_id, name, public_key, created_at FROM public_keys WHERE id = $1`, pubkeyID) | |
| 418 | + | pk := &db.PublicKey{} | |
| 419 | + | err := me.Db.Get(pk, `SELECT id, user_id, name, public_key, created_at FROM public_keys WHERE id = $1`, pubkeyID) | |
| 429 | 420 | if err != nil { | |
| 430 | 421 | return nil, err | |
| 431 | 422 | } | |
| 432 | - | ||
| 433 | - | for rs.Next() { | |
| 434 | - | pk := &db.PublicKey{} | |
| 435 | - | err := rs.Scan(&pk.ID, &pk.UserID, &pk.Name, &pk.Key, &pk.CreatedAt) | |
| 436 | - | if err != nil { | |
| 437 | - | return nil, err | |
| 438 | - | } | |
| 439 | - | ||
| 440 | - | keys = append(keys, pk) | |
| 441 | - | } | |
| 442 | - | ||
| 443 | - | if rs.Err() != nil { | |
| 444 | - | return nil, rs.Err() | |
| 445 | - | } | |
| 446 | - | ||
| 447 | - | if len(keys) == 0 { | |
| 448 | - | return nil, errors.New("no public keys found for key provided") | |
| 449 | - | } | |
| 450 | - | ||
| 451 | - | return keys[0], nil | |
| 423 | + | return pk, nil | |
| 452 | 424 | } | |
| 453 | 425 | ||
| 454 | 426 | func (me *PsqlDB) FindKeysForUser(user *db.User) ([]*db.PublicKey, error) { | |
| 455 | 427 | var keys []*db.PublicKey | |
| 456 | - | rs, err := me.Db.Query(sqlSelectPublicKeys, user.ID) | |
| 428 | + | err := me.Db.Select(&keys, sqlSelectPublicKeys, user.ID) | |
| 457 | 429 | if err != nil { | |
| 458 | - | return keys, err | |
| 459 | - | } | |
| 460 | - | for rs.Next() { | |
| 461 | - | pk := &db.PublicKey{} | |
| 462 | - | err := rs.Scan(&pk.ID, &pk.UserID, &pk.Name, &pk.Key, &pk.CreatedAt) | |
| 463 | - | if err != nil { | |
| 464 | - | return keys, err | |
| 465 | - | } | |
| 466 | - | ||
| 467 | - | keys = append(keys, pk) | |
| 468 | - | } | |
| 469 | - | if rs.Err() != nil { | |
| 470 | - | return keys, rs.Err() | |
| 430 | + | return nil, err | |
| 471 | 431 | } | |
| 472 | 432 | return keys, nil | |
| 473 | 433 | } |
| ... | ... | @@ -559,15 +519,10 @@ func (me *PsqlDB) FindUserByPubkey(key string) (*db.User, error) { | |
| 559 | 519 | ||
| 560 | 520 | func (me *PsqlDB) FindUser(userID string) (*db.User, error) { | |
| 561 | 521 | user := &db.User{} | |
| 562 | - | var un sql.NullString | |
| 563 | - | r := me.Db.QueryRow(sqlSelectUser, userID) | |
| 564 | - | err := r.Scan(&user.ID, &un, &user.CreatedAt) | |
| 522 | + | err := me.Db.Get(user, sqlSelectUser, userID) | |
| 565 | 523 | if err != nil { | |
| 566 | 524 | return nil, err | |
| 567 | 525 | } | |
| 568 | - | if un.Valid { | |
| 569 | - | user.Name = un.String | |
| 570 | - | } | |
| 571 | 526 | return user, nil | |
| 572 | 527 | } | |
| 573 | 528 |
| ... | ... | @@ -589,8 +544,7 @@ func (me *PsqlDB) validateName(name string) (bool, error) { | |
| 589 | 544 | ||
| 590 | 545 | func (me *PsqlDB) FindUserByName(name string) (*db.User, error) { | |
| 591 | 546 | user := &db.User{} | |
| 592 | - | r := me.Db.QueryRow(sqlSelectUserForName, strings.ToLower(name)) | |
| 593 | - | err := r.Scan(&user.ID, &user.Name, &user.CreatedAt) | |
| 547 | + | err := me.Db.Get(user, sqlSelectUserForName, strings.ToLower(name)) | |
| 594 | 548 | if err != nil { | |
| 595 | 549 | return nil, err | |
| 596 | 550 | } |
| ... | ... | @@ -613,13 +567,10 @@ func (me *PsqlDB) findUserForNameAndKey(name string, key string) (*db.User, erro | |
| 613 | 567 | ||
| 614 | 568 | func (me *PsqlDB) FindUserForToken(token string) (*db.User, error) { | |
| 615 | 569 | user := &db.User{} | |
| 616 | - | ||
| 617 | - | r := me.Db.QueryRow(sqlSelectUserForToken, token) | |
| 618 | - | err := r.Scan(&user.ID, &user.Name, &user.CreatedAt) | |
| 570 | + | err := me.Db.Get(user, sqlSelectUserForToken, token) | |
| 619 | 571 | if err != nil { | |
| 620 | 572 | return nil, err | |
| 621 | 573 | } | |
| 622 | - | ||
| 623 | 574 | return user, nil | |
| 624 | 575 | } | |
| 625 | 576 |
| ... | ... | @@ -1126,37 +1077,19 @@ func (me *PsqlDB) FindVisitSiteList(opts *db.SummaryOpts) ([]*db.VisitUrl, error | |
| 1126 | 1077 | ||
| 1127 | 1078 | func (me *PsqlDB) FindUsers() ([]*db.User, error) { | |
| 1128 | 1079 | var users []*db.User | |
| 1129 | - | rs, err := me.Db.Query(sqlSelectUsers) | |
| 1080 | + | err := me.Db.Select(&users, sqlSelectUsers) | |
| 1130 | 1081 | if err != nil { | |
| 1131 | - | return users, err | |
| 1132 | - | } | |
| 1133 | - | for rs.Next() { | |
| 1134 | - | var name sql.NullString | |
| 1135 | - | user := &db.User{} | |
| 1136 | - | err := rs.Scan( | |
| 1137 | - | &user.ID, | |
| 1138 | - | &name, | |
| 1139 | - | &user.CreatedAt, | |
| 1140 | - | ) | |
| 1141 | - | if err != nil { | |
| 1142 | - | return users, err | |
| 1143 | - | } | |
| 1144 | - | user.Name = name.String | |
| 1145 | - | ||
| 1146 | - | users = append(users, user) | |
| 1147 | - | } | |
| 1148 | - | if rs.Err() != nil { | |
| 1149 | - | return users, rs.Err() | |
| 1082 | + | return nil, err | |
| 1150 | 1083 | } | |
| 1151 | 1084 | return users, nil | |
| 1152 | 1085 | } | |
| 1153 | 1086 | ||
| 1154 | - | func (me *PsqlDB) removeTagsForPost(tx *sql.Tx, postID string) error { | |
| 1087 | + | func (me *PsqlDB) removeTagsForPost(tx *sqlx.Tx, postID string) error { | |
| 1155 | 1088 | _, err := tx.Exec(sqlRemoveTagsByPost, postID) | |
| 1156 | 1089 | return err | |
| 1157 | 1090 | } | |
| 1158 | 1091 | ||
| 1159 | - | func (me *PsqlDB) insertTagsForPost(tx *sql.Tx, tags []string, postID string) ([]string, error) { | |
| 1092 | + | func (me *PsqlDB) insertTagsForPost(tx *sqlx.Tx, tags []string, postID string) ([]string, error) { | |
| 1160 | 1093 | ids := make([]string, 0) | |
| 1161 | 1094 | for _, tag := range tags { | |
| 1162 | 1095 | id := "" |
| ... | ... | @@ -1171,13 +1104,12 @@ func (me *PsqlDB) insertTagsForPost(tx *sql.Tx, tags []string, postID string) ([ | |
| 1171 | 1104 | } | |
| 1172 | 1105 | ||
| 1173 | 1106 | func (me *PsqlDB) ReplaceTagsForPost(tags []string, postID string) error { | |
| 1174 | - | ctx := context.Background() | |
| 1175 | - | tx, err := me.Db.BeginTx(ctx, nil) | |
| 1107 | + | tx, err := me.Db.Beginx() | |
| 1176 | 1108 | if err != nil { | |
| 1177 | 1109 | return err | |
| 1178 | 1110 | } | |
| 1179 | 1111 | defer func() { | |
| 1180 | - | err = tx.Rollback() | |
| 1112 | + | _ = tx.Rollback() | |
| 1181 | 1113 | }() | |
| 1182 | 1114 | ||
| 1183 | 1115 | err = me.removeTagsForPost(tx, postID) |
| ... | ... | @@ -1194,12 +1126,12 @@ func (me *PsqlDB) ReplaceTagsForPost(tags []string, postID string) error { | |
| 1194 | 1126 | return err | |
| 1195 | 1127 | } | |
| 1196 | 1128 | ||
| 1197 | - | func (me *PsqlDB) removeAliasesForPost(tx *sql.Tx, postID string) error { | |
| 1129 | + | func (me *PsqlDB) removeAliasesForPost(tx *sqlx.Tx, postID string) error { | |
| 1198 | 1130 | _, err := tx.Exec(sqlRemoveAliasesByPost, postID) | |
| 1199 | 1131 | return err | |
| 1200 | 1132 | } | |
| 1201 | 1133 | ||
| 1202 | - | func (me *PsqlDB) insertAliasesForPost(tx *sql.Tx, aliases []string, postID string) ([]string, error) { | |
| 1134 | + | func (me *PsqlDB) insertAliasesForPost(tx *sqlx.Tx, aliases []string, postID string) ([]string, error) { | |
| 1203 | 1135 | // hardcoded | |
| 1204 | 1136 | denyList := []string{ | |
| 1205 | 1137 | "rss", |
| ... | ... | @@ -1241,13 +1173,12 @@ func (me *PsqlDB) insertAliasesForPost(tx *sql.Tx, aliases []string, postID stri | |
| 1241 | 1173 | } | |
| 1242 | 1174 | ||
| 1243 | 1175 | func (me *PsqlDB) ReplaceAliasesForPost(aliases []string, postID string) error { | |
| 1244 | - | ctx := context.Background() | |
| 1245 | - | tx, err := me.Db.BeginTx(ctx, nil) | |
| 1176 | + | tx, err := me.Db.Beginx() | |
| 1246 | 1177 | if err != nil { | |
| 1247 | 1178 | return err | |
| 1248 | 1179 | } | |
| 1249 | 1180 | defer func() { | |
| 1250 | - | err = tx.Rollback() | |
| 1181 | + | _ = tx.Rollback() | |
| 1251 | 1182 | }() | |
| 1252 | 1183 | ||
| 1253 | 1184 | err = me.removeAliasesForPost(tx, postID) |
| ... | ... | @@ -1342,24 +1273,10 @@ func (me *PsqlDB) FindPopularTags(space string) ([]string, error) { | |
| 1342 | 1273 | ||
| 1343 | 1274 | func (me *PsqlDB) FindFeature(userID string, feature string) (*db.FeatureFlag, error) { | |
| 1344 | 1275 | ff := &db.FeatureFlag{} | |
| 1345 | - | // payment history is allowed to be null | |
| 1346 | - | // https://devtidbits.com/2020/08/03/go-sql-error-converting-null-to-string-is-unsupported/ | |
| 1347 | - | var paymentHistoryID sql.NullString | |
| 1348 | - | err := me.Db.QueryRow(sqlSelectFeatureForUser, userID, feature).Scan( | |
| 1349 | - | &ff.ID, | |
| 1350 | - | &ff.UserID, | |
| 1351 | - | &paymentHistoryID, | |
| 1352 | - | &ff.Name, | |
| 1353 | - | &ff.Data, | |
| 1354 | - | &ff.CreatedAt, | |
| 1355 | - | &ff.ExpiresAt, | |
| 1356 | - | ) | |
| 1276 | + | err := me.Db.Get(ff, sqlSelectFeatureForUser, userID, feature) | |
| 1357 | 1277 | if err != nil { | |
| 1358 | 1278 | return nil, err | |
| 1359 | 1279 | } | |
| 1360 | - | ||
| 1361 | - | ff.PaymentHistoryID = paymentHistoryID | |
| 1362 | - | ||
| 1363 | 1280 | return ff, nil | |
| 1364 | 1281 | } | |
| 1365 | 1282 |
| ... | ... | @@ -1371,31 +1288,9 @@ func (me *PsqlDB) FindFeaturesForUser(userID string) ([]*db.FeatureFlag, error) | |
| 1371 | 1288 | FROM feature_flags | |
| 1372 | 1289 | WHERE user_id=$1 | |
| 1373 | 1290 | ORDER BY name, expires_at DESC;` | |
| 1374 | - | rs, err := me.Db.Query(query, userID) | |
| 1291 | + | err := me.Db.Select(&features, query, userID) | |
| 1375 | 1292 | if err != nil { | |
| 1376 | - | return features, err | |
| 1377 | - | } | |
| 1378 | - | for rs.Next() { | |
| 1379 | - | var paymentHistoryID sql.NullString | |
| 1380 | - | ff := &db.FeatureFlag{} | |
| 1381 | - | err := rs.Scan( | |
| 1382 | - | &ff.ID, | |
| 1383 | - | &ff.UserID, | |
| 1384 | - | &paymentHistoryID, | |
| 1385 | - | &ff.Name, | |
| 1386 | - | &ff.Data, | |
| 1387 | - | &ff.CreatedAt, | |
| 1388 | - | &ff.ExpiresAt, | |
| 1389 | - | ) | |
| 1390 | - | if err != nil { | |
| 1391 | - | return features, err | |
| 1392 | - | } | |
| 1393 | - | ff.PaymentHistoryID = paymentHistoryID | |
| 1394 | - | ||
| 1395 | - | features = append(features, ff) | |
| 1396 | - | } | |
| 1397 | - | if rs.Err() != nil { | |
| 1398 | - | return features, rs.Err() | |
| 1293 | + | return nil, err | |
| 1399 | 1294 | } | |
| 1400 | 1295 | return features, nil | |
| 1401 | 1296 | } |
| ... | ... | @@ -1409,13 +1304,12 @@ func (me *PsqlDB) HasFeatureForUser(userID string, feature string) bool { | |
| 1409 | 1304 | } | |
| 1410 | 1305 | ||
| 1411 | 1306 | func (me *PsqlDB) InsertFeedItems(postID string, items []*db.FeedItem) error { | |
| 1412 | - | ctx := context.Background() | |
| 1413 | - | tx, err := me.Db.BeginTx(ctx, nil) | |
| 1307 | + | tx, err := me.Db.Beginx() | |
| 1414 | 1308 | if err != nil { | |
| 1415 | 1309 | return err | |
| 1416 | 1310 | } | |
| 1417 | 1311 | defer func() { | |
| 1418 | - | err = tx.Rollback() | |
| 1312 | + | _ = tx.Rollback() | |
| 1419 | 1313 | }() | |
| 1420 | 1314 | ||
| 1421 | 1315 | for _, item := range items { |
| ... | ... | @@ -1438,33 +1332,11 @@ func (me *PsqlDB) InsertFeedItems(postID string, items []*db.FeedItem) error { | |
| 1438 | 1332 | } | |
| 1439 | 1333 | ||
| 1440 | 1334 | func (me *PsqlDB) FindFeedItemsByPostID(postID string) ([]*db.FeedItem, error) { | |
| 1441 | - | // sqlSelectFeedItemsByPost | |
| 1442 | - | items := make([]*db.FeedItem, 0) | |
| 1443 | - | rs, err := me.Db.Query(sqlSelectFeedItemsByPost, postID) | |
| 1335 | + | var items []*db.FeedItem | |
| 1336 | + | err := me.Db.Select(&items, sqlSelectFeedItemsByPost, postID) | |
| 1444 | 1337 | if err != nil { | |
| 1445 | - | return items, err | |
| 1446 | - | } | |
| 1447 | - | ||
| 1448 | - | for rs.Next() { | |
| 1449 | - | item := &db.FeedItem{} | |
| 1450 | - | err := rs.Scan( | |
| 1451 | - | &item.ID, | |
| 1452 | - | &item.PostID, | |
| 1453 | - | &item.GUID, | |
| 1454 | - | &item.Data, | |
| 1455 | - | &item.CreatedAt, | |
| 1456 | - | ) | |
| 1457 | - | if err != nil { | |
| 1458 | - | return items, err | |
| 1459 | - | } | |
| 1460 | - | ||
| 1461 | - | items = append(items, item) | |
| 1462 | - | } | |
| 1463 | - | ||
| 1464 | - | if rs.Err() != nil { | |
| 1465 | - | return items, rs.Err() | |
| 1338 | + | return nil, err | |
| 1466 | 1339 | } | |
| 1467 | - | ||
| 1468 | 1340 | return items, nil | |
| 1469 | 1341 | } | |
| 1470 | 1342 |
| ... | ... | @@ -1488,21 +1360,10 @@ func (me *PsqlDB) UpdateProject(userID, name string) error { | |
| 1488 | 1360 | ||
| 1489 | 1361 | func (me *PsqlDB) FindProjectByName(userID, name string) (*db.Project, error) { | |
| 1490 | 1362 | project := &db.Project{} | |
| 1491 | - | r := me.Db.QueryRow(sqlFindProjectByName, userID, name) | |
| 1492 | - | err := r.Scan( | |
| 1493 | - | &project.ID, | |
| 1494 | - | &project.UserID, | |
| 1495 | - | &project.Name, | |
| 1496 | - | &project.ProjectDir, | |
| 1497 | - | &project.Acl, | |
| 1498 | - | &project.Blocked, | |
| 1499 | - | &project.CreatedAt, | |
| 1500 | - | &project.UpdatedAt, | |
| 1501 | - | ) | |
| 1363 | + | err := me.Db.Get(project, sqlFindProjectByName, userID, name) | |
| 1502 | 1364 | if err != nil { | |
| 1503 | 1365 | return nil, err | |
| 1504 | 1366 | } | |
| 1505 | - | ||
| 1506 | 1367 | return project, nil | |
| 1507 | 1368 | } | |
| 1508 | 1369 |
| ... | ... | @@ -1540,24 +1401,12 @@ func (me *PsqlDB) RemoveToken(tokenID string) error { | |
| 1540 | 1401 | } | |
| 1541 | 1402 | ||
| 1542 | 1403 | func (me *PsqlDB) FindTokensForUser(userID string) ([]*db.Token, error) { | |
| 1543 | - | var keys []*db.Token | |
| 1544 | - | rs, err := me.Db.Query(sqlSelectTokensForUser, userID) | |
| 1404 | + | var tokens []*db.Token | |
| 1405 | + | err := me.Db.Select(&tokens, sqlSelectTokensForUser, userID) | |
| 1545 | 1406 | if err != nil { | |
| 1546 | - | return keys, err | |
| 1547 | - | } | |
| 1548 | - | for rs.Next() { | |
| 1549 | - | pk := &db.Token{} | |
| 1550 | - | err := rs.Scan(&pk.ID, &pk.UserID, &pk.Name, &pk.CreatedAt, &pk.ExpiresAt) | |
| 1551 | - | if err != nil { | |
| 1552 | - | return keys, err | |
| 1553 | - | } | |
| 1554 | - | ||
| 1555 | - | keys = append(keys, pk) | |
| 1556 | - | } | |
| 1557 | - | if rs.Err() != nil { | |
| 1558 | - | return keys, rs.Err() | |
| 1407 | + | return nil, err | |
| 1559 | 1408 | } | |
| 1560 | - | return keys, nil | |
| 1409 | + | return tokens, nil | |
| 1561 | 1410 | } | |
| 1562 | 1411 | ||
| 1563 | 1412 | func (me *PsqlDB) InsertFeature(userID, name string, expiresAt time.Time) (*db.FeatureFlag, error) { |
| ... | ... | @@ -1602,13 +1451,12 @@ func (me *PsqlDB) AddPicoPlusUser(username, email, paymentType, txId string) err | |
| 1602 | 1451 | return err | |
| 1603 | 1452 | } | |
| 1604 | 1453 | ||
| 1605 | - | ctx := context.Background() | |
| 1606 | - | tx, err := me.Db.BeginTx(ctx, nil) | |
| 1454 | + | tx, err := me.Db.Beginx() | |
| 1607 | 1455 | if err != nil { | |
| 1608 | 1456 | return err | |
| 1609 | 1457 | } | |
| 1610 | 1458 | defer func() { | |
| 1611 | - | err = tx.Rollback() | |
| 1459 | + | _ = tx.Rollback() | |
| 1612 | 1460 | }() | |
| 1613 | 1461 | ||
| 1614 | 1462 | var paymentHistoryId sql.NullString |
| ... | ... | @@ -1688,60 +1536,24 @@ func (me *PsqlDB) InsertTunsEventLog(log *db.TunsEventLog) error { | |
| 1688 | 1536 | } | |
| 1689 | 1537 | ||
| 1690 | 1538 | func (me *PsqlDB) FindTunsEventLogsByAddr(userID, addr string) ([]*db.TunsEventLog, error) { | |
| 1691 | - | logs := []*db.TunsEventLog{} | |
| 1692 | - | rs, err := me.Db.Query( | |
| 1539 | + | var logs []*db.TunsEventLog | |
| 1540 | + | err := me.Db.Select(&logs, | |
| 1693 | 1541 | `SELECT id, user_id, server_id, remote_addr, event_type, tunnel_type, connection_type, tunnel_id, created_at | |
| 1694 | 1542 | FROM tuns_event_logs WHERE user_id=$1 AND tunnel_id=$2 ORDER BY created_at DESC`, userID, addr) | |
| 1695 | 1543 | if err != nil { | |
| 1696 | 1544 | return nil, err | |
| 1697 | 1545 | } | |
| 1698 | - | ||
| 1699 | - | for rs.Next() { | |
| 1700 | - | log := db.TunsEventLog{} | |
| 1701 | - | err := rs.Scan( | |
| 1702 | - | &log.ID, &log.UserId, &log.ServerID, &log.RemoteAddr, | |
| 1703 | - | &log.EventType, &log.TunnelType, &log.ConnectionType, | |
| 1704 | - | &log.TunnelID, &log.CreatedAt, | |
| 1705 | - | ) | |
| 1706 | - | if err != nil { | |
| 1707 | - | return nil, err | |
| 1708 | - | } | |
| 1709 | - | logs = append(logs, &log) | |
| 1710 | - | } | |
| 1711 | - | ||
| 1712 | - | if rs.Err() != nil { | |
| 1713 | - | return nil, rs.Err() | |
| 1714 | - | } | |
| 1715 | - | ||
| 1716 | 1546 | return logs, nil | |
| 1717 | 1547 | } | |
| 1718 | 1548 | ||
| 1719 | 1549 | func (me *PsqlDB) FindTunsEventLogs(userID string) ([]*db.TunsEventLog, error) { | |
| 1720 | - | logs := []*db.TunsEventLog{} | |
| 1721 | - | rs, err := me.Db.Query( | |
| 1550 | + | var logs []*db.TunsEventLog | |
| 1551 | + | err := me.Db.Select(&logs, | |
| 1722 | 1552 | `SELECT id, user_id, server_id, remote_addr, event_type, tunnel_type, connection_type, tunnel_id, created_at | |
| 1723 | 1553 | FROM tuns_event_logs WHERE user_id=$1 ORDER BY created_at DESC`, userID) | |
| 1724 | 1554 | if err != nil { | |
| 1725 | 1555 | return nil, err | |
| 1726 | 1556 | } | |
| 1727 | - | ||
| 1728 | - | for rs.Next() { | |
| 1729 | - | log := db.TunsEventLog{} | |
| 1730 | - | err := rs.Scan( | |
| 1731 | - | &log.ID, &log.UserId, &log.ServerID, &log.RemoteAddr, | |
| 1732 | - | &log.EventType, &log.TunnelType, &log.ConnectionType, | |
| 1733 | - | &log.TunnelID, &log.CreatedAt, | |
| 1734 | - | ) | |
| 1735 | - | if err != nil { | |
| 1736 | - | return nil, err | |
| 1737 | - | } | |
| 1738 | - | logs = append(logs, &log) | |
| 1739 | - | } | |
| 1740 | - | ||
| 1741 | - | if rs.Err() != nil { | |
| 1742 | - | return nil, rs.Err() | |
| 1743 | - | } | |
| 1744 | - | ||
| 1745 | 1557 | return logs, nil | |
| 1746 | 1558 | } | |
| 1747 | 1559 |
| ... | ... | @@ -1781,79 +1593,29 @@ func (me *PsqlDB) FindUserStats(userID string) (*db.UserStats, error) { | |
| 1781 | 1593 | } | |
| 1782 | 1594 | ||
| 1783 | 1595 | func (me *PsqlDB) FindAccessLogs(userID string, fromDate *time.Time) ([]*db.AccessLog, error) { | |
| 1784 | - | logs := []*db.AccessLog{} | |
| 1785 | - | rs, err := me.Db.Query( | |
| 1786 | - | `SELECT id, user_id, service, pubkey, identity, created_at FROM access_logs WHERE user_id=$1 AND created_at >= $2 ORDER BY created_at DESC`, userID, fromDate) | |
| 1596 | + | var logs []*db.AccessLog | |
| 1597 | + | err := me.Db.Select(&logs, `SELECT id, user_id, service, pubkey, identity, created_at FROM access_logs WHERE user_id=$1 AND created_at >= $2 ORDER BY created_at DESC`, userID, fromDate) | |
| 1787 | 1598 | if err != nil { | |
| 1788 | 1599 | return nil, err | |
| 1789 | 1600 | } | |
| 1790 | - | ||
| 1791 | - | for rs.Next() { | |
| 1792 | - | log := db.AccessLog{} | |
| 1793 | - | err := rs.Scan( | |
| 1794 | - | &log.ID, &log.UserID, &log.Service, &log.Pubkey, &log.Identity, &log.CreatedAt, | |
| 1795 | - | ) | |
| 1796 | - | if err != nil { | |
| 1797 | - | return nil, err | |
| 1798 | - | } | |
| 1799 | - | logs = append(logs, &log) | |
| 1800 | - | } | |
| 1801 | - | ||
| 1802 | - | if rs.Err() != nil { | |
| 1803 | - | return nil, rs.Err() | |
| 1804 | - | } | |
| 1805 | - | ||
| 1806 | 1601 | return logs, nil | |
| 1807 | 1602 | } | |
| 1808 | 1603 | ||
| 1809 | 1604 | func (me *PsqlDB) FindAccessLogsByPubkey(pubkey string, fromDate *time.Time) ([]*db.AccessLog, error) { | |
| 1810 | - | logs := []*db.AccessLog{} | |
| 1811 | - | rs, err := me.Db.Query( | |
| 1812 | - | `SELECT id, user_id, service, pubkey, identity, created_at FROM access_logs WHERE pubkey=$1 AND created_at >= $2 ORDER BY created_at DESC`, pubkey, fromDate) | |
| 1605 | + | var logs []*db.AccessLog | |
| 1606 | + | err := me.Db.Select(&logs, `SELECT id, user_id, service, pubkey, identity, created_at FROM access_logs WHERE pubkey=$1 AND created_at >= $2 ORDER BY created_at DESC`, pubkey, fromDate) | |
| 1813 | 1607 | if err != nil { | |
| 1814 | 1608 | return nil, err | |
| 1815 | 1609 | } | |
| 1816 | - | ||
| 1817 | - | for rs.Next() { | |
| 1818 | - | log := db.AccessLog{} | |
| 1819 | - | err := rs.Scan( | |
| 1820 | - | &log.ID, &log.UserID, &log.Service, &log.Pubkey, &log.Identity, &log.CreatedAt, | |
| 1821 | - | ) | |
| 1822 | - | if err != nil { | |
| 1823 | - | return nil, err | |
| 1824 | - | } | |
| 1825 | - | logs = append(logs, &log) | |
| 1826 | - | } | |
| 1827 | - | ||
| 1828 | - | if rs.Err() != nil { | |
| 1829 | - | return nil, rs.Err() | |
| 1830 | - | } | |
| 1831 | - | ||
| 1832 | 1610 | return logs, nil | |
| 1833 | 1611 | } | |
| 1834 | 1612 | ||
| 1835 | 1613 | func (me *PsqlDB) FindPubkeysInAccessLogs(userID string) ([]string, error) { | |
| 1836 | - | pubkeys := []string{} | |
| 1837 | - | rs, err := me.Db.Query( | |
| 1838 | - | `SELECT DISTINCT(pubkey) FROM access_logs WHERE user_id=$1`, userID, | |
| 1839 | - | ) | |
| 1614 | + | var pubkeys []string | |
| 1615 | + | err := me.Db.Select(&pubkeys, `SELECT DISTINCT(pubkey) FROM access_logs WHERE user_id=$1`, userID) | |
| 1840 | 1616 | if err != nil { | |
| 1841 | 1617 | return nil, err | |
| 1842 | 1618 | } | |
| 1843 | - | ||
| 1844 | - | for rs.Next() { | |
| 1845 | - | pubkey := "" | |
| 1846 | - | err := rs.Scan(&pubkey) | |
| 1847 | - | if err != nil { | |
| 1848 | - | return nil, err | |
| 1849 | - | } | |
| 1850 | - | pubkeys = append(pubkeys, pubkey) | |
| 1851 | - | } | |
| 1852 | - | ||
| 1853 | - | if rs.Err() != nil { | |
| 1854 | - | return nil, rs.Err() | |
| 1855 | - | } | |
| 1856 | - | ||
| 1857 | 1619 | return pubkeys, nil | |
| 1858 | 1620 | } | |
| 1859 | 1621 |
+4
-4
pkg/db/postgres/storage_test.go
#
| ... | ... | @@ -342,7 +342,7 @@ func TestInsertPublicKey_Success(t *testing.T) { | |
| 342 | 342 | t.Fatalf("RegisterUser failed: %v", err) | |
| 343 | 343 | } | |
| 344 | 344 | ||
| 345 | - | err = testDB.InsertPublicKey(user.ID, "ssh-ed25519 AAAAC3NzaC1lZDI1NTE5AAAAI secondkey", "second key", nil) | |
| 345 | + | err = testDB.InsertPublicKey(user.ID, "ssh-ed25519 AAAAC3NzaC1lZDI1NTE5AAAAI secondkey", "second key") | |
| 346 | 346 | if err != nil { | |
| 347 | 347 | t.Fatalf("InsertPublicKey failed: %v", err) | |
| 348 | 348 | } |
| ... | ... | @@ -361,7 +361,7 @@ func TestInsertPublicKey_Duplicate(t *testing.T) { | |
| 361 | 361 | ||
| 362 | 362 | user, _ := testDB.RegisterUser("dupkeyowner", "ssh-ed25519 AAAAC3NzaC1lZDI1NTE5AAAAI dupkey", "comment") | |
| 363 | 363 | ||
| 364 | - | err := testDB.InsertPublicKey(user.ID, "ssh-ed25519 AAAAC3NzaC1lZDI1NTE5AAAAI dupkey", "same key", nil) | |
| 364 | + | err := testDB.InsertPublicKey(user.ID, "ssh-ed25519 AAAAC3NzaC1lZDI1NTE5AAAAI dupkey", "same key") | |
| 365 | 365 | if err == nil { | |
| 366 | 366 | t.Error("expected error for duplicate key, got nil") | |
| 367 | 367 | } |
| ... | ... | @@ -385,7 +385,7 @@ func TestFindKeysForUser(t *testing.T) { | |
| 385 | 385 | cleanupTestData(t) | |
| 386 | 386 | ||
| 387 | 387 | user, _ := testDB.RegisterUser("multikeyowner", "ssh-ed25519 AAAAC3NzaC1lZDI1NTE5AAAAI multikeyowner1", "key1") | |
| 388 | - | _ = testDB.InsertPublicKey(user.ID, "ssh-ed25519 AAAAC3NzaC1lZDI1NTE5AAAAI multikeyowner2", "key2", nil) | |
| 388 | + | _ = testDB.InsertPublicKey(user.ID, "ssh-ed25519 AAAAC3NzaC1lZDI1NTE5AAAAI multikeyowner2", "key2") | |
| 389 | 389 | ||
| 390 | 390 | keys, err := testDB.FindKeysForUser(user) | |
| 391 | 391 | if err != nil { |
| ... | ... | @@ -400,7 +400,7 @@ func TestRemoveKeys(t *testing.T) { | |
| 400 | 400 | cleanupTestData(t) | |
| 401 | 401 | ||
| 402 | 402 | user, _ := testDB.RegisterUser("removekeyowner", "ssh-ed25519 AAAAC3NzaC1lZDI1NTE5AAAAI removekeyowner", "key1") | |
| 403 | - | _ = testDB.InsertPublicKey(user.ID, "ssh-ed25519 AAAAC3NzaC1lZDI1NTE5AAAAI removekeyowner2", "key2", nil) | |
| 403 | + | _ = testDB.InsertPublicKey(user.ID, "ssh-ed25519 AAAAC3NzaC1lZDI1NTE5AAAAI removekeyowner2", "key2") | |
| 404 | 404 | ||
| 405 | 405 | keys, _ := testDB.FindKeysForUser(user) | |
| 406 | 406 | if len(keys) != 2 { |
+1
-2
pkg/db/stub/stub.go
#
| ... | ... | @@ -33,7 +32,7 @@ func (me *StubDB) UpdatePublicKey(pubkeyID, name string) (*db.PublicKey, error) | |
| 33 | 32 | return nil, errNotImpl | |
| 34 | 33 | } | |
| 35 | 34 | ||
| 36 | - | func (me *StubDB) InsertPublicKey(userID, key, name string, tx *sql.Tx) error { | |
| 35 | + | func (me *StubDB) InsertPublicKey(userID, key, name string) error { | |
| 37 | 36 | return errNotImpl | |
| 38 | 37 | } | |
| 39 | 38 |