import sys
import json
import torch
import clip
import psycopg2
from PIL import Image

# device = "cpu" # 搜索量大时建议换成 "cuda"
device = "cuda" if torch.cuda.is_available() else "cpu"

#model, preprocess = clip.load("ViT-B/32", device=device)
model, preprocess = clip.load("ViT-B/32", device=device)
# 1. 提取上传图片的向量
image = preprocess(Image.open(sys.argv[1])).unsqueeze(0).to(device)

with torch.no_grad():
    vec = model.encode_image(image)

vec = vec / vec.norm(dim=-1, keepdim=True)
vector = vec.cpu().numpy()[0].tolist()
vector_str = "[" + ",".join(map(str, vector)) + "]"

# 2. 连接数据库
conn = psycopg2.connect(
    host="192.168.1.42",
    database="postgres",
    user="postgres",
    password="yiparts007"
)
cur = conn.cursor()

# 3. 关联 part_products 表，查询完整的商品信息与相似度
# 使用 <=> (Cosine 相似度) 比 <-> (欧氏距离) 更适合 CLIP 归一化特征向量
cur.execute("""
    SELECT 
        p.brand,
        p.number,
        p.partname_cn,
        p.partname_en,
        p.url,
        i.image_url,
        (1 - (i.embedding <=> %s::vector)) as similarity
    FROM part_images i
    JOIN part_products p ON i.product_id = p.id
    ORDER BY i.embedding <=> %s::vector
    LIMIT 10
""", (vector_str, vector_str))

rows = cur.fetchall()

# 4. 构建统一的 JSON 数据输出结构
results = []
for r in rows:
    results.append({
        "brand": r[0],
        "number": r[1],
        "partname_cn": r[2],
        "partname_en": r[3],
        "url": r[4],
        "image_url": r[5],
        "similarity": round(float(r[6]), 4) # 相似度得分
    })

# 清理连接并向标准输出打印最终结果
cur.close()
conn.close()

print(json.dumps(results, ensure_ascii=True))