package main
import (
"context"
"crypto/sha256"
"encoding/json"
"fmt"
"io"
"net/http"
"os"
"path/filepath"
"time"
"github.com/jackc/pgx/v5"
)
func resetFull(ctx context.Context, c *pgx.Conn) error {
if err := guard(ctx, c); err != nil {
return err
}
_, err := c.Exec(ctx, `DELETE FROM blog_comment WHERE id>500000;
DELETE FROM blog_post_tags WHERE post_id>100000; DELETE FROM blog_post WHERE id>100000;
SELECT setval(pg_get_serial_sequence('blog_post','id'),100000);
SELECT setval(pg_get_serial_sequence('blog_comment','id'),500000);
SELECT setval(pg_get_serial_sequence('blog_post_tags','id'),(SELECT max(id) FROM blog_post_tags));
UPDATE blog_post p SET view_count=s.view_count,updated_at=s.updated_at FROM harness_seed_state s WHERE p.id=s.id`)
return err
}
func fingerprint(ctx context.Context, c *pgx.Conn) (map[string]any, error) {
result := map[string]any{}
for _, table := range []string{"blog_user", "blog_tag", "blog_post", "blog_post_tags", "blog_comment"} {
var count int64
var hash string
err := c.QueryRow(ctx, "SELECT count(*),md5(string_agg(md5(row_to_json(t)::text),'' ORDER BY id)) FROM "+table+" t").Scan(&count, &hash)
if err != nil {
return nil, err
}
result[table] = map[string]any{"count": count, "md5": hash}
}
return result, nil
}
func equivalence(ctx context.Context, referenceDSN, candidateDSN, referenceURL, candidateURL, output string) error {
a, err := openDB(ctx, referenceDSN)
if err != nil {
return err
}
defer a.Close(ctx)
b, err := openDB(ctx, candidateDSN)
if err != nil {
return err
}
defer b.Close(ctx)
if err = resetFull(ctx, a); err != nil {
return err
}
if err = resetFull(ctx, b); err != nil {
return err
}
af, err := fingerprint(ctx, a)
if err != nil {
return err
}
bf, err := fingerprint(ctx, b)
if err != nil {
return err
}
if canonical(af) != canonical(bf) {
return fmt.Errorf("full dataset fingerprints differ before comparison")
}
client := &http.Client{Timeout: 180 * time.Second}
results := []map[string]any{}
for _, path := range []string{"/api/posts", "/api/posts/search?q=python", "/api/posts/by-tag/python", "/api/posts/1", "/api/users/1", "/api/users/find?email=user00000%40example.com"} {
hashes := []string{}
statuses := []int{}
for _, base := range []string{referenceURL, candidateURL} {
start := time.Now()
resp, err := client.Get(base + path)
if err != nil {
return err
}
raw, err := io.ReadAll(resp.Body)
resp.Body.Close()
if err != nil {
return err
}
v, err := decode(raw)
if err != nil {
return err
}
v, err = normalize(v, start, time.Now())
if err != nil {
return err
}
sum := sha256.Sum256([]byte(canonical(v)))
hashes = append(hashes, fmt.Sprintf("%x", sum))
statuses = append(statuses, resp.StatusCode)
}
pass := hashes[0] == hashes[1] && statuses[0] == 200 && statuses[1] == 200
results = append(results, map[string]any{"path": path, "pass": pass, "reference_sha256": hashes[0], "candidate_sha256": hashes[1]})
fmt.Printf("full parity %s pass=%v\n", path, pass)
}
if err = os.MkdirAll(filepath.Dir(output), 0755); err != nil {
return err
}
raw, _ := json.MarshalIndent(map[string]any{"fixture": af, "cases": results}, "", " ")
if err = os.WriteFile(output, append(raw, '\n'), 0644); err != nil {
return err
}
for _, r := range results {
if !r["pass"].(bool) {
return fmt.Errorf("full response mismatch")
}
}
return nil
}