194 lines
6.3 KiB
Go
194 lines
6.3 KiB
Go
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 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) 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) }
|