-
Notifications
You must be signed in to change notification settings - Fork 1
Expand file tree
/
Copy pathnative_model_cache.go
More file actions
117 lines (106 loc) · 3.57 KB
/
Copy pathnative_model_cache.go
File metadata and controls
117 lines (106 loc) · 3.57 KB
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
package main
import (
"bytes"
"context"
"crypto/aes"
"crypto/cipher"
"crypto/hkdf"
"crypto/sha256"
"encoding/base64"
"errors"
"fmt"
)
const qmcV1Prefix = "QMC\x01"
const qmcNonceSize = 12
const qmcTagSize = 16
func deriveModelCacheKey(uid string) ([]byte, error) {
key, err := hkdf.Key(sha256.New, []byte(uid), []byte("qoder-model-cache-enc"), "model-cache-v1", 32)
if err != nil {
return nil, modelCacheBackendFailure(fmt.Errorf("derive model cache key: %w", err))
}
if len(key) != 32 {
return nil, modelCacheBackendFailure(fmt.Errorf("derived model cache key length %d, want 32", len(key)))
}
return key, nil
}
type nativeModelCacheDecryptor struct{}
var _ modelCacheDecryptor = nativeModelCacheDecryptor{}
func (nativeModelCacheDecryptor) Decrypt(ctx context.Context, blob, uid string) ([]byte, error) {
if err := ctx.Err(); err != nil {
return nil, err
}
return decryptQMCV1(blob, uid)
}
func decryptQMCV1(blob, uid string) ([]byte, error) {
envelope, err := decodeStrictStdBase64(blob)
if err != nil {
return nil, modelCacheBackendIncompatible(errors.New("model cache envelope encoding is invalid"))
}
minimumLength := len(qmcV1Prefix) + qmcNonceSize + qmcTagSize
if len(envelope) < minimumLength {
return nil, modelCacheBackendIncompatible(errors.New("model cache envelope is too short"))
}
if !bytes.Equal(envelope[:len(qmcV1Prefix)], []byte(qmcV1Prefix)) {
return nil, modelCacheBackendIncompatible(errors.New("model cache envelope magic or version is unsupported"))
}
nonceStart := len(qmcV1Prefix)
ciphertextStart := nonceStart + qmcNonceSize
nonce := envelope[nonceStart:ciphertextStart]
ciphertext := envelope[ciphertextStart:]
key, err := deriveModelCacheKey(uid)
if err != nil {
return nil, err
}
gcm, err := newQMCGCM(key)
if err != nil {
return nil, err
}
plain, err := gcm.Open(nil, nonce, ciphertext, nil)
if err != nil {
return nil, modelCacheBackendIncompatible(errors.New("model cache envelope authentication failed"))
}
return append([]byte(nil), plain...), nil
}
func encryptQMCV1ForTest(plain []byte, uid string, nonce []byte) (string, error) {
if len(nonce) != qmcNonceSize {
return "", newProtocolError(
protocolInvalidInput,
"model cache input is invalid",
fmt.Errorf("model cache nonce length %d, want %d", len(nonce), qmcNonceSize),
)
}
key, err := deriveModelCacheKey(uid)
if err != nil {
return "", err
}
gcm, err := newQMCGCM(key)
if err != nil {
return "", err
}
sealed := gcm.Seal(nil, nonce, plain, nil)
envelope := make([]byte, 0, len(qmcV1Prefix)+len(nonce)+len(sealed))
envelope = append(envelope, qmcV1Prefix...)
envelope = append(envelope, nonce...)
envelope = append(envelope, sealed...)
return base64.StdEncoding.EncodeToString(envelope), nil
}
func newQMCGCM(key []byte) (cipher.AEAD, error) {
block, err := aes.NewCipher(key)
if err != nil {
return nil, modelCacheBackendFailure(fmt.Errorf("create model cache AES cipher: %w", err))
}
gcm, err := cipher.NewGCM(block)
if err != nil {
return nil, modelCacheBackendFailure(fmt.Errorf("create model cache GCM: %w", err))
}
if gcm.NonceSize() != qmcNonceSize || gcm.Overhead() != qmcTagSize {
return nil, modelCacheBackendFailure(fmt.Errorf("model cache GCM layout nonce/tag lengths %d/%d", gcm.NonceSize(), gcm.Overhead()))
}
return gcm, nil
}
func modelCacheBackendIncompatible(internal error) error {
return newProtocolError(protocolBackendIncompatible, "model cache data is incompatible", internal)
}
func modelCacheBackendFailure(internal error) error {
return newProtocolError(protocolBackendFailure, "model cache operation failed", internal)
}