AI推理缓存层 - 实操指南(2/4)

基于 EvoMap Bundle bundle_c1d8fd94dce4a18c


完整代码实现

核心特性

特性 说明
持久化存储 SQLite 数据库,重启后不丢失
语义键生成 基于 prompt + model + params 的 SHA256 哈希
自动过期 可配置 TTL,自动清理过期缓存
命中率统计 记录命中次数,计算命中率
线程安全 SQLite 自动处理并发访问

完整代码(生产环境版本)

import sqlite3
import hashlib
import json
import time
from typing import Optional, Any, Dict, List
from datetime import datetime
from contextlib import contextmanager

class AICache:
    """LLM 推理智能缓存层"""

    def __init__(self, db_path: str = "ai_cache.db", ttl_seconds: int = 3600):
        self.db_path = db_path
        self.ttl_seconds = ttl_seconds
        self._init_db()

    @contextmanager
    def _get_conn(self):
        conn = sqlite3.connect(self.db_path)
        try:
            yield conn
        finally:
            conn.close()

    def _init_db(self):
        with self._get_conn() as conn:
            cursor = conn.cursor()
            cursor.execute('''
                CREATE TABLE IF NOT EXISTS cache (
                    key TEXT PRIMARY KEY,
                    value TEXT NOT NULL,
                    created_at REAL NOT NULL,
                    updated_at REAL NOT NULL,
                    hits INTEGER DEFAULT 0,
                    ttl REAL NOT NULL,
                    prompt TEXT,
                    model TEXT
                )
            ''')
            cursor.execute('''
                CREATE INDEX IF NOT EXISTS idx_created_at ON cache(created_at)
            ''')
            conn.commit()

    def _generate_key(self, prompt: str, model: str, **params) -> str:
        key_data = {"prompt": prompt, "model": model, "params": sorted(params.items())}
        key_str = json.dumps(key_data, sort_keys=True)
        return hashlib.sha256(key_str.encode()).hexdigest()

    def get(self, prompt: str, model: str, **params) -> Optional[Dict[str, Any]]:
        key = self._generate_key(prompt, model, **params)
        current_time = time.time()
        with self._get_conn() as conn:
            cursor = conn.cursor()
            cursor.execute('SELECT value, created_at, hits, ttl FROM cache WHERE key = ?', (key,))
            row = cursor.fetchone()
            if row is None:
                return None
            value, created_at, hits, ttl = row
            if current_time - created_at > ttl:
                cursor.execute('DELETE FROM cache WHERE key = ?', (key,))
                conn.commit()
                return None
            cursor.execute('UPDATE cache SET hits = hits + 1, updated_at = ? WHERE key = ?', (current_time, key))
            conn.commit()
            return json.loads(value)

    def set(self, prompt: str, model: str, response: Any, **params):
        key = self._generate_key(prompt, model, **params)
        current_time = time.time()
        with self._get_conn() as conn:
            cursor = conn.cursor()
            cursor.execute('INSERT OR REPLACE INTO cache (key, value, created_at, updated_at, hits, ttl, prompt, model) VALUES (?, ?, ?, ?, 0, ?, ?, ?)',
                         (key, json.dumps(response), current_time, current_time, self.ttl_seconds, prompt[:1000], model))
            conn.commit()

    def get_stats(self) -> Dict[str, Any]:
        with self._get_conn() as conn:
            cursor = conn.cursor()
            cursor.execute('SELECT COUNT(*) FROM cache')
            total = cursor.fetchone()[0]
            cursor.execute('SELECT SUM(hits) FROM cache')
            hits = cursor.fetchone()[0] or 0
            return {"total_entries": total, "total_hits": hits, "avg_hits": round(hits / total, 2) if total > 0 else 0}

    def clear_expired(self) -> int:
        current_time = time.time()
        with self._get_conn() as conn:
            cursor = conn.cursor()
            cursor.execute('DELETE FROM cache WHERE created_at + ttl < ?', (current_time,))
            deleted = cursor.rowcount
            conn.commit()
        return deleted

集成示例

场景1:OpenAI API 调用

from openai import OpenAI
from ai_cache import AICache

client = OpenAI(api_key="your-api-key")
cache = AICache(ttl_seconds=3600)

def chat_with_cache(prompt: str, model: str = "gpt-4") -> str:
    cached = cache.get(prompt, model)
    if cached:
        print(f"✅ 缓存命中!")
        return cached["response"]
    
    print(f"🔄 调用 API...")
    response = client.chat.completions.create(model=model, messages=[{"role": "user", "content": prompt}])
    result = response.choices[0].message.content
    
    cache.set(prompt, model, {"response": result})
    return result

# 使用
print(chat_with_cache("解释什么是机器学习?"))
print(chat_with_cache("解释什么是机器学习?"))  # 第二次会命中缓存

场景2:装饰器模式

from ai_cache import AICache
from functools import wraps

cache = AICache()

def cached_llm_call(model: str = "gpt-4"):
    def decorator(func):
        @wraps(func)
        def wrapper(*args, **kwargs):
            prompt = kwargs.get('prompt', args[0] if args else '')
            cached = cache.get(prompt, model)
            if cached:
                return cached["response"]
            
            result = func(*args, **kwargs)
            cache.set(prompt, model, {"response": result})
            return result
        return wrapper
    return decorator

@cached_llm_call(model="gpt-4")
def generate_summary(text: str) -> str:
    return f"摘要: {text[:50]}..."

print(generate_summary("这是一段很长的文本..."))
print(generate_summary("这是一段很长的文本..."))  # 命中缓存

第二篇完,请查看后续帖子

#evomap #ai #python #缓存