#!/usr/bin/env python3 import hashlib import os import sys import time from contextlib import contextmanager from pathlib import Path import psycopg import requests from dotenv import load_dotenv from blacklist import ( EXCLUDE_DIR_NAMES, EXCLUDE_FILENAME_KEYWORDS, EXCLUDE_PATH_PARTS, SENSITIVE_LITERAL_MARKERS, SENSITIVE_REGEX_PATTERNS, ) load_dotenv(Path(__file__).parent.parent.parent / '.env.memory') VAULT_ROOT = Path(os.getenv('VAULT_DIR', '.')).resolve() def require_env(name: str) -> str: value = os.getenv(name, '').strip() if not value: raise RuntimeError(f'缺少必需环境变量: {name}') return value def get_conn(): return psycopg.connect(require_env('PG_DSN')) def normalize_rel(path: Path) -> str: return str(path.resolve().relative_to(VAULT_ROOT)).replace('\\', '/') def safe_rel(path: Path) -> str | None: try: return normalize_rel(path) except Exception: return None def sha256_text(text: str) -> str: return hashlib.sha256(text.encode('utf-8')).hexdigest() def vector_literal(values: list[float]) -> str: return '[' + ','.join(f'{v:.8f}' for v in values) + ']' def _sample_head_mid_tail(content: bytes, span: int = 1200) -> str: size = len(content) if size <= span * 3: return content.decode('utf-8', errors='ignore').lower() mid = max(0, (size // 2) - (span // 2)) sampled = content[:span] + content[mid : mid + span] + content[-span:] return sampled.decode('utf-8', errors='ignore').lower() def is_excluded(file_path: Path) -> str | None: """返回排除原因字符串,未排除则返回 None。""" rel = safe_rel(file_path) if rel is None: return 'path:unresolvable' parts = set(Path(rel).parts) for name in parts.intersection(EXCLUDE_DIR_NAMES): return f'dir:{name}' for name in parts.intersection(EXCLUDE_PATH_PARTS): return f'path:{name}' for kw in EXCLUDE_FILENAME_KEYWORDS: if kw in file_path.name.lower(): return f'filename:{kw}' try: snippet = _sample_head_mid_tail(file_path.read_bytes()) for marker in SENSITIVE_LITERAL_MARKERS: if marker in snippet: return f'content:literal:{marker[:20]}' for pattern in SENSITIVE_REGEX_PATTERNS: if pattern.search(snippet): return f'content:regex:{pattern.pattern[:30]}' except Exception: return 'content:read_error' return None @contextmanager def index_lock(): lock_file = Path(os.getenv('INDEX_LOCK_FILE', str(VAULT_ROOT / '.memory-index.lock'))) lock_file.parent.mkdir(parents=True, exist_ok=True) with open(lock_file, 'w', encoding='utf-8') as fh: if sys.platform == 'win32': import msvcrt msvcrt.locking(fh.fileno(), msvcrt.LK_LOCK, 1) try: yield finally: msvcrt.locking(fh.fileno(), msvcrt.LK_UNLCK, 1) else: import fcntl as _fcntl _fcntl.flock(fh, _fcntl.LOCK_EX) try: yield finally: _fcntl.flock(fh, _fcntl.LOCK_UN) def embed_text(text: str) -> list[float]: api_key = require_env('OPENROUTER_API_KEY') base_url = os.getenv('OPENROUTER_BASE_URL', 'https://openrouter.ai/api/v1').rstrip('/') model = require_env('OPENROUTER_EMBED_MODEL') expected_dim = int(os.getenv('OPENROUTER_EMBED_DIM', '1536')) headers = { 'Authorization': f'Bearer {api_key}', 'Content-Type': 'application/json', } referer = os.getenv('OPENROUTER_HTTP_REFERER', '').strip() if referer: headers['HTTP-Referer'] = referer title = os.getenv('OPENROUTER_X_OPENROUTER_TITLE', '').strip() or os.getenv( 'OPENROUTER_X_TITLE', '' ).strip() if title: headers['X-OpenRouter-Title'] = title payload = {'model': model, 'input': text, 'encoding_format': 'float'} last_error: Exception | None = None for attempt in range(1, 4): try: response = requests.post( f'{base_url}/embeddings', headers=headers, json=payload, timeout=30, ) response.raise_for_status() body = response.json() data = body.get('data') if not isinstance(data, list) or not data: raise RuntimeError(f'embedding 响应缺少 data 字段: keys={list(body.keys())}') embedding = data[0].get('embedding') if not isinstance(embedding, list): raise RuntimeError('embedding 响应结构异常: data[0].embedding 缺失') if len(embedding) != expected_dim: raise RuntimeError( f'embedding 维度不匹配: got={len(embedding)} expected={expected_dim}' ) return embedding except Exception as error: last_error = error if attempt < 3: time.sleep(0.8 * attempt) raise RuntimeError(f'OpenRouter embedding 请求失败: {last_error}') from last_error