import sys
import json
import torch
import clip
import psycopg2
from PIL import Image

# 0. 参数校验
if len(sys.argv) < 2:
    print(json.dumps({"error": "Please provide an image path."}))
    sys.exit(1)

# 自动选择设备
device = "cuda" if torch.cuda.is_available() else "cpu"

try:
    # 1. 加载模型（建议先用 ViT-B/32 探路，显存占用低）
    # 如果显存空闲了，可以换成 "ViT-L/14"
    #model, preprocess = clip.load("ViT-B/32", device=device)
    model, preprocess = clip.load("RN50x64", device=device)

    # 提取并处理上传图片的向量
    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)) + "]"

    # 【重要优化】算完向量后，立刻释放 GPU 显存，不给共用服务器添堵
    if device == "cuda":
        del model, image, vec
        torch.cuda.empty_cache()

except Exception as e:
    print(json.dumps({"error": f"Model inference failed: {str(e)}"}))
    sys.exit(1)


# 2. 连接数据库
try:
    conn = psycopg2.connect(
        host="192.168.1.42",
        database="postgres",
        user="postgres",
        password="yiparts007",
        connect_timeout=5 # 增加超时控制
    )
    cur = conn.cursor()

    # 3. 关联查询完整的商品信息与相似度
    # 使用 <=> (Cosine 相似度距离)
    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_1024 i
        JOIN part_products_1024 p ON i.product_id = p.id
        ORDER BY i.embedding <=> %s::vector
        LIMIT 4
    """, (vector_str, vector_str))

    rows = cur.fetchall()

    # 4. 构建统一的 JSON 数据输出结构
    results = []
    if rows:
        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) if r[6] is not None else 0.0
            })

    # 清理连接
    cur.close()
    conn.close()

    # 向标准输出打印最终结果
    print(json.dumps(results, ensure_ascii=False)) # ensure_ascii=False 可以让中文不乱码

except Exception as e:
    print(json.dumps({"error": f"Database operation failed: {str(e)}"}))
    sys.exit(1)