mirror of
https://github.com/Tencent/WeKnora.git
synced 2026-09-24 16:29:01 +08:00
refactor: replace CryptoService with utils AES helpers and simplify encryption
- Remove infrastructure/crypto package (CryptoService, config, PBKDF2-based encryption) - Replace with lightweight utils.GetAESKey/EncryptAESGCM/DecryptAESGCM helpers - Remove cryptoSvc dependency from modelService, tenantService, weKnoraCloudService - Remove crypto state persistence logic from container initialization - Add TruncatePromptTokens field to WeKnoraCloud embed request - Update frontend hint text for WeKnoraCloud credential errors Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Opus 4.6
parent
2c0c1b8bab
commit
d4edd374df
@@ -28,7 +28,7 @@
|
||||
<t-icon name="error-circle" style="font-size: 16px; color: #f97316; flex-shrink: 0; margin-top: 1px;" />
|
||||
<div style="font-size: 13px; color: #9a3412; line-height: 1.5;">
|
||||
<strong>WeKnoraCloud 凭证已失效</strong><br />
|
||||
{{ weKnoraCloudReinitReason || '服务重启后加密密钥已变更,已保存的凭证无法解密。' }}请重新填写 APPID 和 APPSECRET 并点击"保存并初始化"以恢复服务。
|
||||
{{ weKnoraCloudReinitReason || '服务重启后加密密钥已变更,已保存的凭证无法解密。' }}
|
||||
</div>
|
||||
</div>
|
||||
<div class="weknoracloud-config-card" style="background: var(--td-bg-color-container); border: 1px solid var(--td-component-stroke); border-radius: 8px; padding: 20px;">
|
||||
|
||||
@@ -4,7 +4,6 @@ import (
|
||||
"context"
|
||||
"errors"
|
||||
|
||||
"github.com/Tencent/WeKnora/internal/infrastructure/crypto"
|
||||
"github.com/Tencent/WeKnora/internal/logger"
|
||||
"github.com/Tencent/WeKnora/internal/models/asr"
|
||||
"github.com/Tencent/WeKnora/internal/models/chat"
|
||||
@@ -14,6 +13,7 @@ import (
|
||||
"github.com/Tencent/WeKnora/internal/models/vlm"
|
||||
"github.com/Tencent/WeKnora/internal/types"
|
||||
"github.com/Tencent/WeKnora/internal/types/interfaces"
|
||||
"github.com/Tencent/WeKnora/internal/utils"
|
||||
)
|
||||
|
||||
// ErrModelNotFound is returned when a model cannot be found in the repository
|
||||
@@ -24,20 +24,17 @@ type modelService struct {
|
||||
repo interfaces.ModelRepository
|
||||
ollamaService *ollama.OllamaService
|
||||
pooler embedding.EmbedderPooler
|
||||
cryptoSvc *crypto.CryptoService
|
||||
}
|
||||
|
||||
// NewModelService creates a new model service instance
|
||||
func NewModelService(repo interfaces.ModelRepository,
|
||||
ollamaService *ollama.OllamaService,
|
||||
pooler embedding.EmbedderPooler,
|
||||
cryptoSvc *crypto.CryptoService,
|
||||
) interfaces.ModelService {
|
||||
return &modelService{
|
||||
repo: repo,
|
||||
ollamaService: ollamaService,
|
||||
pooler: pooler,
|
||||
cryptoSvc: cryptoSvc,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -46,11 +43,12 @@ func (s *modelService) decryptAppSecret(encrypted string) string {
|
||||
if encrypted == "" {
|
||||
return encrypted
|
||||
}
|
||||
plain, err := s.cryptoSvc.DecryptString(encrypted)
|
||||
if err != nil {
|
||||
return ""
|
||||
if key := utils.GetAESKey(); key != nil {
|
||||
if encrypted, err := utils.DecryptAESGCM(encrypted, key); err == nil {
|
||||
return encrypted
|
||||
}
|
||||
}
|
||||
return plain
|
||||
return encrypted
|
||||
}
|
||||
|
||||
// CreateModel creates a new model in the repository
|
||||
|
||||
@@ -13,7 +13,6 @@ import (
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/Tencent/WeKnora/internal/infrastructure/crypto"
|
||||
"github.com/Tencent/WeKnora/internal/logger"
|
||||
"github.com/Tencent/WeKnora/internal/types"
|
||||
"github.com/Tencent/WeKnora/internal/types/interfaces"
|
||||
@@ -34,13 +33,12 @@ type ListTenantsParams struct {
|
||||
|
||||
// tenantService implements the TenantService interface
|
||||
type tenantService struct {
|
||||
repo interfaces.TenantRepository // Repository for tenant data operations
|
||||
cryptoSvc *crypto.CryptoService
|
||||
repo interfaces.TenantRepository // Repository for tenant data operations
|
||||
}
|
||||
|
||||
// NewTenantService creates a new tenant service instance
|
||||
func NewTenantService(repo interfaces.TenantRepository, cryptoSvc *crypto.CryptoService) interfaces.TenantService {
|
||||
return &tenantService{repo: repo, cryptoSvc: cryptoSvc}
|
||||
func NewTenantService(repo interfaces.TenantRepository) interfaces.TenantService {
|
||||
return &tenantService{repo: repo}
|
||||
}
|
||||
|
||||
// CreateTenant creates a new tenant
|
||||
@@ -353,11 +351,12 @@ func (s *tenantService) GetDocreaderCredentials(ctx context.Context) *types.Docr
|
||||
if appID == "" || tenant.ParserEngineConfig.DocreaderAPIKey == "" {
|
||||
return nil
|
||||
}
|
||||
apiKey, err := s.cryptoSvc.DecryptString(tenant.ParserEngineConfig.DocreaderAPIKey)
|
||||
if err != nil || apiKey == "" {
|
||||
return nil
|
||||
if key := utils.GetAESKey(); key != nil {
|
||||
if encrypted, err := utils.DecryptAESGCM(tenant.ParserEngineConfig.DocreaderAPIKey, key); err == nil {
|
||||
return &types.DocreaderCredentials{AppID: appID, APIKey: encrypted}
|
||||
}
|
||||
}
|
||||
return &types.DocreaderCredentials{AppID: appID, APIKey: apiKey}
|
||||
return &types.DocreaderCredentials{AppID: appID, APIKey: tenant.ParserEngineConfig.DocreaderAPIKey}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -372,10 +371,12 @@ func (s *tenantService) GetDocreaderCredentials(ctx context.Context) *types.Docr
|
||||
if err == nil && tenant != nil && tenant.ParserEngineConfig != nil {
|
||||
appID := strings.TrimSpace(tenant.ParserEngineConfig.DocreaderAppID)
|
||||
if appID != "" && tenant.ParserEngineConfig.DocreaderAPIKey != "" {
|
||||
apiKey, err := s.cryptoSvc.DecryptString(tenant.ParserEngineConfig.DocreaderAPIKey)
|
||||
if err == nil && apiKey != "" {
|
||||
return &types.DocreaderCredentials{AppID: appID, APIKey: apiKey}
|
||||
if key := utils.GetAESKey(); key != nil {
|
||||
if encrypted, err := utils.DecryptAESGCM(tenant.ParserEngineConfig.DocreaderAPIKey, key); err == nil {
|
||||
return &types.DocreaderCredentials{AppID: appID, APIKey: encrypted}
|
||||
}
|
||||
}
|
||||
return &types.DocreaderCredentials{AppID: appID, APIKey: tenant.ParserEngineConfig.DocreaderAPIKey}
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -7,29 +7,28 @@ import (
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/Tencent/WeKnora/internal/infrastructure/crypto"
|
||||
"github.com/Tencent/WeKnora/internal/models/provider"
|
||||
"github.com/Tencent/WeKnora/internal/types"
|
||||
"github.com/Tencent/WeKnora/internal/types/interfaces"
|
||||
"github.com/Tencent/WeKnora/internal/utils"
|
||||
"github.com/google/uuid"
|
||||
)
|
||||
|
||||
const WeKnoraCloudAPI = "https://weknora.woa.com/platform/openapi"
|
||||
|
||||
type weKnoraCloudService struct {
|
||||
repo interfaces.ModelRepository
|
||||
tenantRepo interfaces.TenantRepository
|
||||
cryptoSvc *crypto.CryptoService
|
||||
}
|
||||
|
||||
// NewWeKnoraCloudService 构造 WeKnoraCloudService
|
||||
func NewWeKnoraCloudService(
|
||||
repo interfaces.ModelRepository,
|
||||
tenantRepo interfaces.TenantRepository,
|
||||
cryptoSvc *crypto.CryptoService,
|
||||
) interfaces.WeKnoraCloudService {
|
||||
return &weKnoraCloudService{
|
||||
repo: repo,
|
||||
tenantRepo: tenantRepo,
|
||||
cryptoSvc: cryptoSvc,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -99,9 +98,11 @@ func (s *weKnoraCloudService) prepareInitialize(ctx context.Context, appID, appS
|
||||
}
|
||||
|
||||
// Step 2: Encrypt appSecret
|
||||
encryptedSecret, err := s.cryptoSvc.EncryptString(appSecret)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("encrypt AppSecret failed: %w", err)
|
||||
encryptedSecret := appSecret
|
||||
if key := utils.GetAESKey(); key != nil {
|
||||
if encrypted, err := utils.EncryptAESGCM(appSecret, key); err == nil {
|
||||
encryptedSecret = encrypted
|
||||
}
|
||||
}
|
||||
|
||||
// Step 3: Get tenantID and take snapshot
|
||||
@@ -296,18 +297,19 @@ func (s *weKnoraCloudService) CheckStatus(ctx context.Context) (*types.WeKnoraCl
|
||||
return &types.WeKnoraCloudStatusResult{
|
||||
HasModels: true,
|
||||
NeedsReinit: true,
|
||||
Reason: "WeKnoraCloud 凭证为空,请重新填写 APPID 和 API Key",
|
||||
Reason: fmt.Sprintf("WeKnoraCloud 凭证为空,请重新填写 APPID 和 API Key, 请前往:%s", WeKnoraCloudAPI),
|
||||
}, nil
|
||||
}
|
||||
|
||||
// Try to decrypt the API key
|
||||
_, decErr := s.cryptoSvc.DecryptString(tenant.ParserEngineConfig.DocreaderAPIKey)
|
||||
if decErr != nil {
|
||||
return &types.WeKnoraCloudStatusResult{
|
||||
HasModels: true,
|
||||
NeedsReinit: true,
|
||||
Reason: "WeKnoraCloud 凭证解密失败(服务重启后加密密钥已变更),请重新填写 APPID 和 API Key",
|
||||
}, nil
|
||||
if key := utils.GetAESKey(); key != nil {
|
||||
if _, err := utils.DecryptAESGCM(tenant.ParserEngineConfig.DocreaderAPIKey, key); err != nil {
|
||||
return &types.WeKnoraCloudStatusResult{
|
||||
HasModels: true,
|
||||
NeedsReinit: true,
|
||||
Reason: "WeKnoraCloud 凭证解密失败(服务重启后加密密钥已变更),请重新填写 APPID 和 API Key",
|
||||
}, nil
|
||||
}
|
||||
}
|
||||
|
||||
return &types.WeKnoraCloudStatusResult{HasModels: true, NeedsReinit: false}, nil
|
||||
|
||||
@@ -6,7 +6,6 @@ package container
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"net/url"
|
||||
@@ -64,7 +63,6 @@ import (
|
||||
"github.com/Tencent/WeKnora/internal/im/telegram"
|
||||
"github.com/Tencent/WeKnora/internal/im/wechat"
|
||||
"github.com/Tencent/WeKnora/internal/im/wecom"
|
||||
"github.com/Tencent/WeKnora/internal/infrastructure/crypto"
|
||||
"github.com/Tencent/WeKnora/internal/infrastructure/docparser"
|
||||
infra_web_search "github.com/Tencent/WeKnora/internal/infrastructure/web_search"
|
||||
"github.com/Tencent/WeKnora/internal/logger"
|
||||
@@ -100,7 +98,6 @@ func BuildContainer(container *dig.Container) *dig.Container {
|
||||
// Core infrastructure configuration
|
||||
logger.Debugf(ctx, "[Container] Registering core infrastructure...")
|
||||
must(container.Provide(config.LoadConfig))
|
||||
must(container.Provide(initCryptoService))
|
||||
must(container.Provide(initTracer))
|
||||
must(container.Provide(initDatabase))
|
||||
must(container.Provide(initFileService))
|
||||
@@ -311,92 +308,6 @@ func initTracer() (*tracing.Tracer, error) {
|
||||
return tracing.InitTracer()
|
||||
}
|
||||
|
||||
// cryptoStateFile 是持久化 crypto 状态的文件路径,存储在 data-files volume 中
|
||||
// 保证重启后能恢复相同的 masterKey 和 salt,避免已加密数据无法解密
|
||||
const cryptoStateFile = "/data/files/.crypto_state.json"
|
||||
|
||||
// cryptoState 持久化存储的 crypto 配置
|
||||
type cryptoState struct {
|
||||
MasterKey string `json:"master_key"`
|
||||
Salt string `json:"salt"` // base64 编码
|
||||
}
|
||||
|
||||
func initCryptoService() (*crypto.CryptoService, error) {
|
||||
cfg := crypto.LoadConfigFromEnv()
|
||||
|
||||
// 如果环境变量同时提供了 MasterKey 和 Salt,直接使用(优先级最高)
|
||||
if cfg.MasterKey != "" && cfg.Salt != "" {
|
||||
return crypto.NewCryptoServiceFromConfig(cfg)
|
||||
}
|
||||
|
||||
// 尝试从持久化文件恢复状态
|
||||
if state, err := loadCryptoState(cryptoStateFile); err == nil {
|
||||
// 文件读取成功,环境变量中未覆盖的字段从文件补充
|
||||
if cfg.MasterKey == "" {
|
||||
cfg.MasterKey = state.MasterKey
|
||||
}
|
||||
if cfg.Salt == "" {
|
||||
cfg.Salt = state.Salt
|
||||
}
|
||||
return crypto.NewCryptoServiceFromConfig(cfg)
|
||||
}
|
||||
|
||||
// 文件不存在或读取失败,使用/生成配置后写入文件
|
||||
if cfg.MasterKey == "" {
|
||||
cfg.MasterKey = "weknora-default-key"
|
||||
}
|
||||
|
||||
// 调用 NewCryptoServiceFromConfig 会在 cfg.Salt 为空时自动生成随机 salt
|
||||
svc, err := crypto.NewCryptoServiceFromConfig(cfg)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// 将实际使用的 masterKey 和 salt 持久化,供重启恢复
|
||||
saltB64 := base64.StdEncoding.EncodeToString(svc.GetSalt())
|
||||
state := &cryptoState{
|
||||
MasterKey: cfg.MasterKey,
|
||||
Salt: saltB64,
|
||||
}
|
||||
if saveErr := saveCryptoState(cryptoStateFile, state); saveErr != nil {
|
||||
// 持久化失败仅记录警告,不阻止服务启动
|
||||
logger.Warnf(context.Background(), "[CryptoService] 无法持久化 crypto 状态到 %s: %v (重启后已加密数据将无法解密)", cryptoStateFile, saveErr)
|
||||
} else {
|
||||
logger.Infof(context.Background(), "[CryptoService] crypto 状态已持久化到 %s", cryptoStateFile)
|
||||
}
|
||||
|
||||
return svc, nil
|
||||
}
|
||||
|
||||
// loadCryptoState 从文件加载 crypto 状态
|
||||
func loadCryptoState(path string) (*cryptoState, error) {
|
||||
data, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
var state cryptoState
|
||||
if err := json.Unmarshal(data, &state); err != nil {
|
||||
return nil, fmt.Errorf("parse crypto state: %w", err)
|
||||
}
|
||||
if state.MasterKey == "" || state.Salt == "" {
|
||||
return nil, fmt.Errorf("invalid crypto state: missing fields")
|
||||
}
|
||||
return &state, nil
|
||||
}
|
||||
|
||||
// saveCryptoState 将 crypto 状态写入文件
|
||||
func saveCryptoState(path string, state *cryptoState) error {
|
||||
if err := os.MkdirAll(filepath.Dir(path), 0o700); err != nil {
|
||||
return fmt.Errorf("mkdir: %w", err)
|
||||
}
|
||||
data, err := json.Marshal(state)
|
||||
if err != nil {
|
||||
return fmt.Errorf("marshal: %w", err)
|
||||
}
|
||||
// 权限 0o600:仅 owner 可读写
|
||||
return os.WriteFile(path, data, 0o600)
|
||||
}
|
||||
|
||||
func initRedisClient() (*redis.Client, error) {
|
||||
redisAddr := os.Getenv("REDIS_ADDR")
|
||||
if redisAddr == "" {
|
||||
|
||||
@@ -1,224 +0,0 @@
|
||||
package crypto
|
||||
|
||||
import (
|
||||
"encoding/base64"
|
||||
"fmt"
|
||||
"os"
|
||||
"strconv"
|
||||
)
|
||||
|
||||
// Config 加密服务配置结构
|
||||
type Config struct {
|
||||
// MasterKey 主密钥,用于派生加密密钥
|
||||
// 可以从环境变量或配置文件中读取
|
||||
MasterKey string `json:"master_key" yaml:"master_key" env:"CRYPTO_MASTER_KEY"`
|
||||
|
||||
// Salt 盐值,用于密钥派生
|
||||
// 如果为空,将自动生成随机盐值
|
||||
Salt string `json:"salt" yaml:"salt" env:"CRYPTO_SALT"`
|
||||
|
||||
// SaltLength 盐值长度(当Salt为空时使用)
|
||||
SaltLength int `json:"salt_length" yaml:"salt_length" env:"CRYPTO_SALT_LENGTH" default:"16"`
|
||||
|
||||
// KeyDerivationIterations PBKDF2迭代次数
|
||||
KeyDerivationIterations int `json:"key_derivation_iterations" yaml:"key_derivation_iterations" env:"CRYPTO_ITERATIONS" default:"10000"`
|
||||
|
||||
// KeyLength 派生密钥长度
|
||||
KeyLength int `json:"key_length" yaml:"key_length" env:"CRYPTO_KEY_LENGTH" default:"32"`
|
||||
}
|
||||
|
||||
// DefaultConfig 返回默认配置
|
||||
func DefaultConfig() *Config {
|
||||
return &Config{
|
||||
MasterKey: "",
|
||||
Salt: "",
|
||||
SaltLength: 16,
|
||||
KeyDerivationIterations: 10000,
|
||||
KeyLength: 32,
|
||||
}
|
||||
}
|
||||
|
||||
// LoadConfigFromEnv 从环境变量加载配置
|
||||
func LoadConfigFromEnv() *Config {
|
||||
config := DefaultConfig()
|
||||
|
||||
// 从环境变量读取主密钥
|
||||
if masterKey := os.Getenv("CRYPTO_MASTER_KEY"); masterKey != "" {
|
||||
config.MasterKey = masterKey
|
||||
}
|
||||
|
||||
// 从环境变量读取盐值
|
||||
if salt := os.Getenv("CRYPTO_SALT"); salt != "" {
|
||||
config.Salt = salt
|
||||
}
|
||||
|
||||
// 从环境变量读取盐值长度
|
||||
if saltLengthStr := os.Getenv("CRYPTO_SALT_LENGTH"); saltLengthStr != "" {
|
||||
if saltLength, err := strconv.Atoi(saltLengthStr); err == nil && saltLength >= 8 {
|
||||
config.SaltLength = saltLength
|
||||
}
|
||||
}
|
||||
|
||||
// 从环境变量读取迭代次数
|
||||
if iterationsStr := os.Getenv("CRYPTO_ITERATIONS"); iterationsStr != "" {
|
||||
if iterations, err := strconv.Atoi(iterationsStr); err == nil && iterations > 0 {
|
||||
config.KeyDerivationIterations = iterations
|
||||
}
|
||||
}
|
||||
|
||||
// 从环境变量读取密钥长度
|
||||
if keyLengthStr := os.Getenv("CRYPTO_KEY_LENGTH"); keyLengthStr != "" {
|
||||
if keyLength, err := strconv.Atoi(keyLengthStr); err == nil && keyLength >= 16 {
|
||||
config.KeyLength = keyLength
|
||||
}
|
||||
}
|
||||
|
||||
return config
|
||||
}
|
||||
|
||||
// Validate 验证配置是否有效
|
||||
func (c *Config) Validate() error {
|
||||
if c.MasterKey == "" {
|
||||
return fmt.Errorf("master key cannot be empty")
|
||||
}
|
||||
|
||||
if c.SaltLength < 8 {
|
||||
return fmt.Errorf("salt length must be at least 8 bytes")
|
||||
}
|
||||
|
||||
if c.KeyDerivationIterations < 1000 {
|
||||
return fmt.Errorf("key derivation iterations must be at least 1000")
|
||||
}
|
||||
|
||||
if c.KeyLength < 16 {
|
||||
return fmt.Errorf("key length must be at least 16 bytes")
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetSaltBytes 获取盐值的字节数组形式
|
||||
func (c *Config) GetSaltBytes() ([]byte, error) {
|
||||
if c.Salt != "" {
|
||||
// 如果配置了盐值,解码Base64格式
|
||||
saltBytes, err := base64.StdEncoding.DecodeString(c.Salt)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to decode salt from base64: %w", err)
|
||||
}
|
||||
return saltBytes, nil
|
||||
}
|
||||
|
||||
// 如果没有配置盐值,生成随机盐值
|
||||
salt, err := GenerateRandomSalt(c.SaltLength)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to generate random salt: %w", err)
|
||||
}
|
||||
|
||||
return salt, nil
|
||||
}
|
||||
|
||||
// NewCryptoServiceFromConfig 从配置创建加密服务
|
||||
func NewCryptoServiceFromConfig(config *Config) (*CryptoService, error) {
|
||||
if err := config.Validate(); err != nil {
|
||||
return nil, fmt.Errorf("invalid config: %w", err)
|
||||
}
|
||||
|
||||
saltBytes, err := config.GetSaltBytes()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to get salt bytes: %w", err)
|
||||
}
|
||||
|
||||
return NewCryptoService(config.MasterKey, saltBytes)
|
||||
}
|
||||
|
||||
// ExampleConfigUsage 展示配置使用示例
|
||||
func ExampleConfigUsage() {
|
||||
fmt.Println("=== 配置使用示例 ===")
|
||||
|
||||
// 方法1: 手动创建配置
|
||||
manualConfig := &Config{
|
||||
MasterKey: "my-secret-master-key",
|
||||
Salt: "", // 自动生成盐值
|
||||
SaltLength: 16,
|
||||
KeyDerivationIterations: 10000,
|
||||
KeyLength: 32,
|
||||
}
|
||||
|
||||
cryptoService1, err := NewCryptoServiceFromConfig(manualConfig)
|
||||
if err != nil {
|
||||
fmt.Printf("手动配置创建失败: %v\n", err)
|
||||
} else {
|
||||
fmt.Println("手动配置创建成功")
|
||||
_ = cryptoService1
|
||||
}
|
||||
|
||||
// 方法2: 从环境变量加载配置
|
||||
// 首先设置环境变量
|
||||
os.Setenv("CRYPTO_MASTER_KEY", "env-master-key")
|
||||
os.Setenv("CRYPTO_SALT", base64.StdEncoding.EncodeToString([]byte("env-salt-123456")))
|
||||
os.Setenv("CRYPTO_SALT_LENGTH", "16")
|
||||
os.Setenv("CRYPTO_ITERATIONS", "10000")
|
||||
os.Setenv("CRYPTO_KEY_LENGTH", "32")
|
||||
|
||||
envConfig := LoadConfigFromEnv()
|
||||
cryptoService2, err := NewCryptoServiceFromConfig(envConfig)
|
||||
if err != nil {
|
||||
fmt.Printf("环境变量配置创建失败: %v\n", err)
|
||||
} else {
|
||||
fmt.Println("环境变量配置创建成功")
|
||||
_ = cryptoService2
|
||||
}
|
||||
|
||||
// 清理环境变量
|
||||
os.Unsetenv("CRYPTO_MASTER_KEY")
|
||||
os.Unsetenv("CRYPTO_SALT")
|
||||
os.Unsetenv("CRYPTO_SALT_LENGTH")
|
||||
os.Unsetenv("CRYPTO_ITERATIONS")
|
||||
os.Unsetenv("CRYPTO_KEY_LENGTH")
|
||||
}
|
||||
|
||||
// GenerateConfigExample 生成配置示例
|
||||
func GenerateConfigExample() {
|
||||
fmt.Println("\n=== 配置生成示例 ===")
|
||||
|
||||
// 生成随机主密钥
|
||||
masterKey, err := GenerateRandomKey(32)
|
||||
if err != nil {
|
||||
fmt.Printf("生成主密钥失败: %v\n", err)
|
||||
return
|
||||
}
|
||||
|
||||
// 生成随机盐值
|
||||
salt, err := GenerateRandomSalt(16)
|
||||
if err != nil {
|
||||
fmt.Printf("生成盐值失败: %v\n", err)
|
||||
return
|
||||
}
|
||||
|
||||
saltBase64 := base64.StdEncoding.EncodeToString(salt)
|
||||
|
||||
fmt.Println("生成的配置示例:")
|
||||
fmt.Printf("CRYPTO_MASTER_KEY=%s\n", masterKey)
|
||||
fmt.Printf("CRYPTO_SALT=%s\n", saltBase64)
|
||||
fmt.Printf("CRYPTO_SALT_LENGTH=16\n")
|
||||
fmt.Printf("CRYPTO_ITERATIONS=10000\n")
|
||||
fmt.Printf("CRYPTO_KEY_LENGTH=32\n")
|
||||
|
||||
fmt.Println("\nYAML格式配置示例:")
|
||||
fmt.Printf(`master_key: %s
|
||||
salt: %s
|
||||
salt_length: 16
|
||||
key_derivation_iterations: 10000
|
||||
key_length: 32
|
||||
`, masterKey, saltBase64)
|
||||
|
||||
fmt.Println("\nJSON格式配置示例:")
|
||||
fmt.Printf(`{
|
||||
"master_key": "%s",
|
||||
"salt": "%s",
|
||||
"salt_length": 16,
|
||||
"key_derivation_iterations": 10000,
|
||||
"key_length": 32
|
||||
}
|
||||
`, masterKey, saltBase64)
|
||||
}
|
||||
@@ -1,163 +0,0 @@
|
||||
package crypto
|
||||
|
||||
import (
|
||||
"crypto/aes"
|
||||
"crypto/cipher"
|
||||
"crypto/rand"
|
||||
"crypto/sha256"
|
||||
"encoding/base64"
|
||||
"errors"
|
||||
"fmt"
|
||||
|
||||
"golang.org/x/crypto/pbkdf2"
|
||||
)
|
||||
|
||||
// CryptoService 提供密钥加密解密服务
|
||||
type CryptoService struct {
|
||||
masterKey []byte
|
||||
salt []byte
|
||||
}
|
||||
|
||||
// NewCryptoService 创建新的加密服务实例
|
||||
// masterKey: 主密钥,用于派生加密密钥
|
||||
// salt: 盐值,用于密钥派生,如果为空则生成随机盐
|
||||
func NewCryptoService(masterKey string, salt []byte) (*CryptoService, error) {
|
||||
if masterKey == "" {
|
||||
return nil, errors.New("master key cannot be empty")
|
||||
}
|
||||
|
||||
// 如果未提供盐值,生成随机盐
|
||||
if salt == nil || len(salt) == 0 {
|
||||
salt = make([]byte, 16)
|
||||
if _, err := rand.Read(salt); err != nil {
|
||||
return nil, fmt.Errorf("failed to generate salt: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
return &CryptoService{
|
||||
masterKey: []byte(masterKey),
|
||||
salt: salt,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// deriveKey 使用PBKDF2派生加密密钥
|
||||
func (cs *CryptoService) deriveKey() []byte {
|
||||
return pbkdf2.Key(cs.masterKey, cs.salt, 10000, 32, sha256.New)
|
||||
}
|
||||
|
||||
// Encrypt 加密数据
|
||||
func (cs *CryptoService) Encrypt(plaintext []byte) (string, error) {
|
||||
if len(plaintext) == 0 {
|
||||
return "", errors.New("plaintext cannot be empty")
|
||||
}
|
||||
|
||||
key := cs.deriveKey()
|
||||
|
||||
// 创建AES块密码
|
||||
block, err := aes.NewCipher(key)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("failed to create cipher: %w", err)
|
||||
}
|
||||
|
||||
// 生成随机IV
|
||||
iv := make([]byte, aes.BlockSize)
|
||||
if _, err := rand.Read(iv); err != nil {
|
||||
return "", fmt.Errorf("failed to generate IV: %w", err)
|
||||
}
|
||||
|
||||
// 使用CTR模式加密
|
||||
stream := cipher.NewCTR(block, iv)
|
||||
ciphertext := make([]byte, len(plaintext))
|
||||
stream.XORKeyStream(ciphertext, plaintext)
|
||||
|
||||
// 组合IV和密文
|
||||
result := make([]byte, len(iv)+len(ciphertext))
|
||||
copy(result[:aes.BlockSize], iv)
|
||||
copy(result[aes.BlockSize:], ciphertext)
|
||||
|
||||
// 返回Base64编码的结果
|
||||
return base64.StdEncoding.EncodeToString(result), nil
|
||||
}
|
||||
|
||||
// Decrypt 解密数据
|
||||
func (cs *CryptoService) Decrypt(encryptedData string) ([]byte, error) {
|
||||
if encryptedData == "" {
|
||||
return nil, errors.New("encrypted data cannot be empty")
|
||||
}
|
||||
|
||||
// 解码Base64数据
|
||||
data, err := base64.StdEncoding.DecodeString(encryptedData)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to decode base64: %w", err)
|
||||
}
|
||||
|
||||
if len(data) < aes.BlockSize {
|
||||
return nil, errors.New("encrypted data too short")
|
||||
}
|
||||
|
||||
key := cs.deriveKey()
|
||||
|
||||
// 创建AES块密码
|
||||
block, err := aes.NewCipher(key)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create cipher: %w", err)
|
||||
}
|
||||
|
||||
// 提取IV和密文
|
||||
iv := data[:aes.BlockSize]
|
||||
ciphertext := data[aes.BlockSize:]
|
||||
|
||||
// 使用CTR模式解密
|
||||
stream := cipher.NewCTR(block, iv)
|
||||
plaintext := make([]byte, len(ciphertext))
|
||||
stream.XORKeyStream(plaintext, ciphertext)
|
||||
|
||||
return plaintext, nil
|
||||
}
|
||||
|
||||
// EncryptString 加密字符串
|
||||
func (cs *CryptoService) EncryptString(plaintext string) (string, error) {
|
||||
return cs.Encrypt([]byte(plaintext))
|
||||
}
|
||||
|
||||
// DecryptString 解密字符串
|
||||
func (cs *CryptoService) DecryptString(encryptedData string) (string, error) {
|
||||
plaintext, err := cs.Decrypt(encryptedData)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return string(plaintext), nil
|
||||
}
|
||||
|
||||
// GetSalt 获取当前使用的盐值
|
||||
func (cs *CryptoService) GetSalt() []byte {
|
||||
return cs.salt
|
||||
}
|
||||
|
||||
// GenerateRandomKey 生成随机密钥
|
||||
func GenerateRandomKey(length int) (string, error) {
|
||||
if length < 16 {
|
||||
return "", errors.New("key length must be at least 16 bytes")
|
||||
}
|
||||
|
||||
key := make([]byte, length)
|
||||
if _, err := rand.Read(key); err != nil {
|
||||
return "", fmt.Errorf("failed to generate random key: %w", err)
|
||||
}
|
||||
|
||||
return base64.StdEncoding.EncodeToString(key), nil
|
||||
}
|
||||
|
||||
// GenerateRandomSalt 生成随机盐值
|
||||
func GenerateRandomSalt(length int) ([]byte, error) {
|
||||
if length < 8 {
|
||||
return nil, errors.New("salt length must be at least 8 bytes")
|
||||
}
|
||||
|
||||
salt := make([]byte, length)
|
||||
if _, err := rand.Read(salt); err != nil {
|
||||
return nil, fmt.Errorf("failed to generate random salt: %w", err)
|
||||
}
|
||||
|
||||
return salt, nil
|
||||
}
|
||||
@@ -54,8 +54,9 @@ func NewWeKnoraCloudEmbedder(config Config) (*WeKnoraCloudEmbedder, error) {
|
||||
}
|
||||
|
||||
type weKnoraCloudEmbedRequest struct {
|
||||
Model string `json:"model"`
|
||||
Input []string `json:"input"`
|
||||
Model string `json:"model"`
|
||||
Input []string `json:"input"`
|
||||
TruncatePromptTokens int `json:"truncate_prompt_tokens,omitempty"`
|
||||
}
|
||||
|
||||
type weKnoraCloudEmbedResponse struct {
|
||||
@@ -77,7 +78,7 @@ func (e *WeKnoraCloudEmbedder) Embed(ctx context.Context, text string) ([]float3
|
||||
}
|
||||
|
||||
func (e *WeKnoraCloudEmbedder) BatchEmbed(ctx context.Context, texts []string) ([][]float32, error) {
|
||||
reqBody := weKnoraCloudEmbedRequest{Model: e.effectiveModelName(), Input: texts}
|
||||
reqBody := weKnoraCloudEmbedRequest{Model: e.effectiveModelName(), Input: texts, TruncatePromptTokens: 512}
|
||||
bodyBytes, err := json.Marshal(reqBody)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("weknoracloud embedder: marshal: %w", err)
|
||||
|
||||
Reference in New Issue
Block a user