package postgres import ( "context" "database/sql" "encoding/json" "fmt" "os" "time" "git.nianxx.cn/wangxuming/NianAIGC/backend/internal/assets" ) const assetFields = `a.id, a.owner_id, a.kind, a.name, a.url, a.storage_path, a.source, a.tags, a.metadata, a.created_at, a.updated_at` const ListOwnerAssetsSQL = `SELECT ` + assetFields + ` FROM public.assets AS a WHERE a.owner_id = $1::text ORDER BY a.created_at DESC` const GetOwnerAssetSQL = `SELECT ` + assetFields + ` FROM public.assets AS a WHERE a.owner_id = $1::text AND a.id = $2::text LIMIT 1` const GetOwnerAssetByStoragePathSQL = `SELECT ` + assetFields + ` FROM public.assets AS a WHERE a.owner_id = $1::text AND a.storage_path = $2::text LIMIT 1` const listPublicAccessibleAssetIDs = `SELECT unnest(j.input_asset_ids || j.output_asset_ids) AS asset_id FROM ( SELECT input_asset_ids, output_asset_ids FROM public.generation_jobs WHERE owner_id = $1::text AND external_client_id = $3::text ORDER BY created_at DESC LIMIT $4::integer ) AS j` const ListPublicAssetsSQL = `SELECT ` + assetFields + ` FROM public.assets AS a WHERE a.owner_id = $1::text AND ($2::text = ANY(a.tags) OR a.id IN (` + listPublicAccessibleAssetIDs + `)) ORDER BY a.created_at DESC` const getPublicAccessibleAssetIDs = `SELECT unnest(j.input_asset_ids || j.output_asset_ids) AS asset_id FROM ( SELECT input_asset_ids, output_asset_ids FROM public.generation_jobs WHERE owner_id = $1::text AND external_client_id = $4::text ORDER BY created_at DESC LIMIT $5::integer ) AS j` const GetPublicAssetSQL = `SELECT ` + assetFields + ` FROM public.assets AS a WHERE a.owner_id = $1::text AND a.id = $2::text AND ($3::text = ANY(a.tags) OR a.id IN (` + getPublicAccessibleAssetIDs + `)) LIMIT 1` const CreateAssetSQL = `INSERT INTO public.assets ( id, owner_id, kind, name, url, storage_path, source, tags, metadata, created_at, updated_at ) VALUES ($1::text, $2::text, $3::text, $4::text, $5::text, $6::text, $7::text, $8::text[], $9::jsonb, $10::timestamptz, $11::timestamptz) RETURNING id, owner_id, kind, name, url, storage_path, source, tags, metadata, created_at, updated_at` const DeleteOwnerAssetSQL = `DELETE FROM public.assets WHERE owner_id = $1::text AND id = $2::text RETURNING id, owner_id, kind, name, url, storage_path, source, tags, metadata, created_at, updated_at` func (db *Database) ListOwner(ctx context.Context, owner string) ([]assets.Asset, error) { return db.listAssets(ctx, ListOwnerAssetsSQL, owner) } func (db *Database) GetOwner(ctx context.Context, owner, id string) (assets.Asset, bool, error) { return db.oneAsset(ctx, GetOwnerAssetSQL, owner, id) } func (db *Database) GetOwnerByStoragePath(ctx context.Context, owner, storagePath string) (assets.Asset, bool, error) { return db.oneAsset(ctx, GetOwnerAssetByStoragePathSQL, owner, storagePath) } func (db *Database) ListPublic(ctx context.Context, owner, client string, limit int) ([]assets.Asset, error) { return db.listAssets(ctx, ListPublicAssetsSQL, owner, assets.ClientTag(client), client, limit) } func (db *Database) GetPublic(ctx context.Context, owner, client, id string, limit int) (assets.Asset, bool, error) { return db.oneAsset(ctx, GetPublicAssetSQL, owner, id, assets.ClientTag(client), client, limit) } func (db *Database) Create(ctx context.Context, a assets.Asset) (assets.Asset, error) { if err := db.available(); err != nil { return assets.Asset{}, err } if a.Tags == nil { a.Tags = []string{} } if a.Metadata == nil { a.Metadata = map[string]any{} } metadata, err := json.Marshal(a.Metadata) if err != nil { return assets.Asset{}, fmt.Errorf("encode asset metadata: %w", err) } return db.requiredAsset(ctx, CreateAssetSQL, a.ID, a.OwnerID, string(a.Kind), a.Name, a.URL, optionalDatabaseText(a.StoragePath), string(a.Source), a.Tags, metadata, a.CreatedAt, a.UpdatedAt) } func (db *Database) DeleteOwner(ctx context.Context, owner, id string) (assets.Asset, bool, error) { return db.oneAsset(ctx, DeleteOwnerAssetSQL, owner, id) } func (db *Database) available() error { if db.config.Backend != BackendPostgres || db.querier == nil { return fmt.Errorf("PostgreSQL is unavailable when ZHINIAN_DATA_BACKEND=%s", db.config.Backend) } return nil } func (db *Database) listAssets(ctx context.Context, query string, args ...any) ([]assets.Asset, error) { if err := db.available(); err != nil { return nil, err } rows, err := db.querier.Query(ctx, query, args...) if err != nil { return nil, fmt.Errorf("query assets: %w", err) } defer rows.Close() out := []assets.Asset{} for rows.Next() { a, err := scanAsset(rows) if err != nil { return nil, err } out = append(out, a) } if err = rows.Err(); err != nil { return nil, fmt.Errorf("read assets: %w", err) } return out, nil } func (db *Database) oneAsset(ctx context.Context, query string, args ...any) (assets.Asset, bool, error) { if err := db.available(); err != nil { return assets.Asset{}, false, err } rows, err := db.querier.Query(ctx, query, args...) if err != nil { return assets.Asset{}, false, fmt.Errorf("query asset: %w", err) } defer rows.Close() if !rows.Next() { if err = rows.Err(); err != nil { return assets.Asset{}, false, fmt.Errorf("read asset: %w", err) } return assets.Asset{}, false, nil } a, err := scanAsset(rows) if err != nil { return assets.Asset{}, false, err } if err = rows.Err(); err != nil { return assets.Asset{}, false, fmt.Errorf("read asset: %w", err) } return a, true, nil } func (db *Database) requiredAsset(ctx context.Context, query string, args ...any) (assets.Asset, error) { a, found, err := db.oneAsset(ctx, query, args...) if err != nil { return assets.Asset{}, err } if !found { return assets.Asset{}, fmt.Errorf("asset write returned no row") } return a, nil } func scanAsset(rows Rows) (assets.Asset, error) { var a assets.Asset var kind, source string var storage sql.NullString var metadata []byte var created, updated time.Time if err := rows.Scan(&a.ID, &a.OwnerID, &kind, &a.Name, &a.URL, &storage, &source, &a.Tags, &metadata, &created, &updated); err != nil { return assets.Asset{}, fmt.Errorf("scan asset: %w", err) } a.Kind = assets.Kind(kind) a.Source = assets.Source(source) if storage.Valid { a.StoragePath = storage.String } a.CreatedAt = created a.UpdatedAt = updated if len(metadata) == 0 { metadata = []byte(`{}`) } if err := json.Unmarshal(metadata, &a.Metadata); err != nil { return assets.Asset{}, fmt.Errorf("decode asset metadata: %w", err) } if a.Tags == nil { a.Tags = []string{} } if a.Metadata == nil { a.Metadata = map[string]any{} } return a, nil } func readContractFixture(path string) ([]byte, error) { return os.ReadFile(path) }