File size: 8,456 Bytes
b30d305
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
"""
API密钥加密工具
使用AES加密确保数据库中API密钥的安全性
"""
import os
import base64
import secrets
from typing import Optional, Tuple
from cryptography.fernet import Fernet
from cryptography.hazmat.primitives import hashes
from cryptography.hazmat.primitives.kdf.pbkdf2 import PBKDF2HMAC

from src.utils.logger import setup_logger
from src.utils.env_config import env_config

logger = setup_logger("encryption")


class APIKeyEncryption:
    """API密钥加密管理器"""
    
    def __init__(self):
        self._fernet = None
        self._init_encryption_key()
    
    def _init_encryption_key(self):
        """初始化加密密钥"""
        # 尝试从环境变量获取加密密钥
        encryption_key = os.getenv('ENCRYPTION_KEY')
        
        if not encryption_key:
            # 首先尝试从数据库获取已存储的密钥
            encryption_key = self._get_stored_encryption_key()
            
            if not encryption_key:
                # 检查数据库中是否已存在加密数据
                has_encrypted_data = self._check_existing_encrypted_data()
                
                if has_encrypted_data:
                    logger.error("Found encrypted API keys in database but no ENCRYPTION_KEY available!")
                    logger.error("Option 1: Set ENCRYPTION_KEY in your .env file if you have the key")
                    logger.error("Option 2: Delete encrypted channels and restart to generate new key")
                    raise ValueError("Missing ENCRYPTION_KEY - cannot decrypt existing encrypted data")
                else:
                    # 没有加密数据,生成新密钥并自动保存到数据库配置
                    encryption_key = self._generate_encryption_key()
                    self._store_encryption_key(encryption_key)
                    logger.info("Generated new encryption key and stored in database")
                    logger.info("For better security, consider moving this to .env file:")
                    logger.info(f"ENCRYPTION_KEY={encryption_key}")
            else:
                logger.info("Using stored encryption key from database")
        
        try:
            # 验证密钥格式
            self._fernet = Fernet(encryption_key.encode())
            logger.info("Encryption system initialized successfully")
        except Exception as e:
            logger.error(f"Failed to initialize encryption: {e}")
            raise ValueError("Invalid encryption key format")
    
    def _check_existing_encrypted_data(self) -> bool:
        """检查数据库中是否存在加密数据(避免循环导入)"""
        try:
            from src.utils.env_config import env_config
            import sqlite3
            
            db_path = env_config.database_path
            if not os.path.exists(db_path):
                return False
                
            conn = sqlite3.connect(db_path)
            cursor = conn.execute(
                "SELECT COUNT(*) FROM channels WHERE api_key LIKE 'encrypted:%'"
            )
            count = cursor.fetchone()[0]
            conn.close()
            return count > 0
        except Exception:
            # 如果查询失败(表不存在等),假设没有加密数据
            return False
    
    def _generate_encryption_key(self) -> str:
        """生成新的加密密钥"""
        # 生成32字节的随机密钥
        key = Fernet.generate_key()
        return key.decode()
    
    def _store_encryption_key(self, encryption_key: str):
        """将加密密钥存储到数据库配置中"""
        try:
            from src.utils.env_config import env_config
            import sqlite3
            
            db_path = env_config.database_path
            os.makedirs(os.path.dirname(db_path), exist_ok=True)
            
            conn = sqlite3.connect(db_path)
            # 创建配置表(如果不存在)
            conn.execute('''
                CREATE TABLE IF NOT EXISTS config (
                    key TEXT PRIMARY KEY,
                    value TEXT NOT NULL,
                    created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
                    updated_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP
                )
            ''')
            
            # 存储加密密钥
            conn.execute(
                'INSERT OR REPLACE INTO config (key, value, updated_at) VALUES (?, ?, CURRENT_TIMESTAMP)',
                ('encryption_key', encryption_key)
            )
            conn.commit()
            conn.close()
        except Exception as e:
            logger.warning(f"Failed to store encryption key in database: {e}")
    
    def _get_stored_encryption_key(self) -> Optional[str]:
        """从数据库配置中获取加密密钥"""
        try:
            from src.utils.env_config import env_config
            import sqlite3
            
            db_path = env_config.database_path
            if not os.path.exists(db_path):
                return None
                
            conn = sqlite3.connect(db_path)
            cursor = conn.execute('SELECT value FROM config WHERE key = ?', ('encryption_key',))
            result = cursor.fetchone()
            conn.close()
            
            return result[0] if result else None
        except Exception:
            return None
    
    def encrypt_api_key(self, api_key: str) -> str:
        """加密API密钥"""
        if not api_key:
            return ""
        
        try:
            encrypted_data = self._fernet.encrypt(api_key.encode())
            # 返回base64编码的加密数据,添加前缀标识
            return f"encrypted:{base64.b64encode(encrypted_data).decode()}"
        except Exception as e:
            logger.error(f"Failed to encrypt API key: {e}")
            raise ValueError("Encryption failed")
    
    def decrypt_api_key(self, encrypted_api_key: str) -> str:
        """解密API密钥"""
        if not encrypted_api_key:
            return ""
        
        # 检查是否是加密格式
        if not encrypted_api_key.startswith("encrypted:"):
            # 兼容未加密的旧数据
            logger.warning("Found unencrypted API key, consider re-saving to encrypt it")
            return encrypted_api_key
        
        try:
            # 移除前缀并解码
            encrypted_data = encrypted_api_key[10:]  # 移除 "encrypted:" 前缀
            encrypted_bytes = base64.b64decode(encrypted_data.encode())
            
            # 解密
            decrypted_data = self._fernet.decrypt(encrypted_bytes)
            return decrypted_data.decode()
        except Exception as e:
            logger.error(f"Failed to decrypt API key: {e}")
            raise ValueError("Decryption failed - possibly wrong encryption key")
    
    def is_encrypted(self, data: str) -> bool:
        """检查数据是否已加密"""
        return data.startswith("encrypted:") if data else False
    
    def rotate_encryption_key(self, new_key: str, old_encrypted_data: list) -> list:
        """
        轮换加密密钥(高级功能)
        重新加密所有数据使用新密钥
        """
        # 保存当前密钥
        old_fernet = self._fernet
        
        try:
            # 设置新密钥
            self._fernet = Fernet(new_key.encode())
            
            # 重新加密所有数据
            reencrypted_data = []
            for encrypted_item in old_encrypted_data:
                if self.is_encrypted(encrypted_item):
                    # 使用旧密钥解密
                    self._fernet = old_fernet
                    decrypted = self.decrypt_api_key(encrypted_item)
                    
                    # 使用新密钥加密
                    self._fernet = Fernet(new_key.encode())
                    reencrypted = self.encrypt_api_key(decrypted)
                    reencrypted_data.append(reencrypted)
                else:
                    reencrypted_data.append(encrypted_item)
            
            logger.info(f"Successfully rotated encryption key for {len(reencrypted_data)} items")
            return reencrypted_data
            
        except Exception as e:
            # 恢复旧密钥
            self._fernet = old_fernet
            logger.error(f"Key rotation failed: {e}")
            raise


# 全局加密管理器实例
encryption_manager = APIKeyEncryption()