111 lines
4.0 KiB
Go
111 lines
4.0 KiB
Go
package postgres
|
|
|
|
import (
|
|
"context"
|
|
"encoding/json"
|
|
"fmt"
|
|
|
|
"git.nianxx.cn/wangxuming/NianAIGC/backend/internal/layercompositions"
|
|
)
|
|
|
|
const FindLayerCompositionSQL = `SELECT job_id, owner_id, base_asset_id, base_visible, layers, version, created_at, updated_at
|
|
FROM public.seedream_layer_compositions
|
|
WHERE owner_id = $1::text AND job_id = $2::text
|
|
LIMIT 1`
|
|
|
|
const SaveLayerCompositionSQL = `WITH updated AS (
|
|
UPDATE public.seedream_layer_compositions SET
|
|
base_asset_id = $3::text,
|
|
base_visible = $4::boolean,
|
|
layers = $5::jsonb,
|
|
version = seedream_layer_compositions.version + 1,
|
|
updated_at = now()
|
|
WHERE job_id = $1::text
|
|
AND owner_id = $2::text
|
|
AND version = $6::bigint
|
|
AND $6::bigint > 0
|
|
RETURNING job_id, owner_id, base_asset_id, base_visible, layers, version, created_at, updated_at
|
|
), inserted AS (
|
|
INSERT INTO public.seedream_layer_compositions (
|
|
job_id, owner_id, base_asset_id, base_visible, layers, version, created_at, updated_at
|
|
)
|
|
SELECT $1::text, $2::text, $3::text, $4::boolean, $5::jsonb, 1, now(), now()
|
|
WHERE $6::bigint = 0
|
|
ON CONFLICT (job_id) DO NOTHING
|
|
RETURNING job_id, owner_id, base_asset_id, base_visible, layers, version, created_at, updated_at
|
|
)
|
|
SELECT job_id, owner_id, base_asset_id, base_visible, layers, version, created_at, updated_at FROM updated
|
|
UNION ALL
|
|
SELECT job_id, owner_id, base_asset_id, base_visible, layers, version, created_at, updated_at FROM inserted
|
|
LIMIT 1`
|
|
|
|
func (db *Database) FindLayerComposition(ctx context.Context, ownerID, jobID string) (layercompositions.Composition, bool, error) {
|
|
if err := db.available(); err != nil {
|
|
return layercompositions.Composition{}, false, err
|
|
}
|
|
rows, err := db.querier.Query(ctx, FindLayerCompositionSQL, ownerID, jobID)
|
|
if err != nil {
|
|
return layercompositions.Composition{}, false, fmt.Errorf("find layer composition: %w", err)
|
|
}
|
|
defer rows.Close()
|
|
if !rows.Next() {
|
|
if err := rows.Err(); err != nil {
|
|
return layercompositions.Composition{}, false, fmt.Errorf("read layer composition: %w", err)
|
|
}
|
|
return layercompositions.Composition{}, false, nil
|
|
}
|
|
composition, err := scanLayerComposition(rows)
|
|
if err != nil {
|
|
return layercompositions.Composition{}, false, err
|
|
}
|
|
return composition, true, rows.Err()
|
|
}
|
|
|
|
func (db *Database) SaveLayerComposition(ctx context.Context, composition layercompositions.Composition, expectedVersion int64) (layercompositions.Composition, error) {
|
|
if err := db.available(); err != nil {
|
|
return layercompositions.Composition{}, err
|
|
}
|
|
layers, err := json.Marshal(composition.Layers)
|
|
if err != nil {
|
|
return layercompositions.Composition{}, fmt.Errorf("encode layer composition: %w", err)
|
|
}
|
|
rows, err := db.querier.Query(ctx, SaveLayerCompositionSQL,
|
|
composition.JobID, composition.OwnerID, composition.BaseAssetID,
|
|
composition.BaseVisible, layers, expectedVersion,
|
|
)
|
|
if err != nil {
|
|
return layercompositions.Composition{}, fmt.Errorf("save layer composition: %w", err)
|
|
}
|
|
defer rows.Close()
|
|
if !rows.Next() {
|
|
if err := rows.Err(); err != nil {
|
|
return layercompositions.Composition{}, fmt.Errorf("read saved layer composition: %w", err)
|
|
}
|
|
return layercompositions.Composition{}, layercompositions.ErrConflict
|
|
}
|
|
saved, err := scanLayerComposition(rows)
|
|
if err != nil {
|
|
return layercompositions.Composition{}, err
|
|
}
|
|
return saved, rows.Err()
|
|
}
|
|
|
|
func scanLayerComposition(rows Rows) (layercompositions.Composition, error) {
|
|
var composition layercompositions.Composition
|
|
var rawLayers []byte
|
|
if err := rows.Scan(
|
|
&composition.JobID, &composition.OwnerID, &composition.BaseAssetID,
|
|
&composition.BaseVisible, &rawLayers, &composition.Version,
|
|
&composition.CreatedAt, &composition.UpdatedAt,
|
|
); err != nil {
|
|
return layercompositions.Composition{}, fmt.Errorf("scan layer composition: %w", err)
|
|
}
|
|
if err := json.Unmarshal(rawLayers, &composition.Layers); err != nil {
|
|
return layercompositions.Composition{}, fmt.Errorf("decode layer composition: %w", err)
|
|
}
|
|
if composition.Layers == nil {
|
|
composition.Layers = []layercompositions.Layer{}
|
|
}
|
|
return composition, nil
|
|
}
|