chore: refresh gramsrv public release
This commit is contained in:
parent
75cebe8dbf
commit
70b6820474
1274 changed files with 378751 additions and 59919 deletions
282
internal/app/files/appearance_seed.go
Normal file
282
internal/app/files/appearance_seed.go
Normal file
|
|
@ -0,0 +1,282 @@
|
|||
package files
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"fmt"
|
||||
"hash"
|
||||
|
||||
"telesrv/internal/domain"
|
||||
"telesrv/internal/seed/appearance"
|
||||
)
|
||||
|
||||
// AppearanceSeedStats reports default appearance resources imported into media storage.
|
||||
type AppearanceSeedStats struct {
|
||||
Wallpapers int
|
||||
Documents int
|
||||
Blobs int
|
||||
Skipped bool
|
||||
}
|
||||
|
||||
// SeedAppearance imports the default wallpaper document catalog into telesrv media storage.
|
||||
func (s *Service) SeedAppearance(ctx context.Context) (AppearanceSeedStats, error) {
|
||||
var stats AppearanceSeedStats
|
||||
catalog := appearance.Default()
|
||||
if len(catalog.Wallpapers) == 0 && len(catalog.ChatThemes) == 0 {
|
||||
stats.Skipped = true
|
||||
return stats, nil
|
||||
}
|
||||
stateHash, err := s.seedAppearanceStateHash()
|
||||
if err != nil {
|
||||
return stats, err
|
||||
}
|
||||
ready, err := s.appearanceSeedReady(ctx, catalog)
|
||||
if err != nil {
|
||||
return stats, err
|
||||
}
|
||||
matched, err := s.seedStateMatches(ctx, seedAppearanceStateKey, stateHash)
|
||||
if err != nil {
|
||||
return stats, err
|
||||
}
|
||||
if matched && ready {
|
||||
stats.Wallpapers = len(catalog.Wallpapers)
|
||||
stats.Skipped = true
|
||||
return stats, nil
|
||||
}
|
||||
seen := make(map[int64]bool)
|
||||
seedDoc := func(in appearance.Document, label string) error {
|
||||
if in.ID == 0 || seen[in.ID] {
|
||||
return nil
|
||||
}
|
||||
seen[in.ID] = true
|
||||
doc, blobs, err := s.seedAppearanceDocument(ctx, in)
|
||||
if err != nil {
|
||||
return fmt.Errorf("seed %s %d: %w", label, in.ID, err)
|
||||
}
|
||||
if doc.ID != 0 {
|
||||
stats.Documents++
|
||||
}
|
||||
stats.Blobs += blobs
|
||||
return nil
|
||||
}
|
||||
for _, wallpaper := range catalog.Wallpapers {
|
||||
if err := seedDoc(wallpaper.Document, "wallpaper"); err != nil {
|
||||
return stats, err
|
||||
}
|
||||
stats.Wallpapers++
|
||||
}
|
||||
// 聊天主题的主题背景墙纸文档也要进媒体库,否则客户端取主题背景会 404。
|
||||
for _, ct := range catalog.ChatThemes {
|
||||
for _, setting := range ct.Settings {
|
||||
if err := seedDoc(setting.Wallpaper.Document, "chat theme wallpaper"); err != nil {
|
||||
return stats, err
|
||||
}
|
||||
}
|
||||
}
|
||||
if err := s.putSeedState(ctx, seedAppearanceStateKey, stateHash); err != nil {
|
||||
return stats, err
|
||||
}
|
||||
return stats, nil
|
||||
}
|
||||
|
||||
func (s *Service) seedAppearanceStateHash() (string, error) {
|
||||
raw, err := appearance.FS.ReadFile("default_appearance_seed.json")
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return seedStateHash(func(h hash.Hash) error {
|
||||
writeSeedStateHeader(h, seedAppearanceStateVersion, s.dc)
|
||||
_, _ = h.Write(raw)
|
||||
return nil
|
||||
})
|
||||
}
|
||||
|
||||
func (s *Service) appearanceSeedReady(ctx context.Context, catalog appearance.Catalog) (bool, error) {
|
||||
docs := make(map[int64]appearance.Document)
|
||||
add := func(doc appearance.Document) {
|
||||
if doc.ID != 0 {
|
||||
docs[doc.ID] = doc
|
||||
}
|
||||
}
|
||||
for _, wallpaper := range catalog.Wallpapers {
|
||||
add(wallpaper.Document)
|
||||
}
|
||||
for _, ct := range catalog.ChatThemes {
|
||||
for _, setting := range ct.Settings {
|
||||
add(setting.Wallpaper.Document)
|
||||
}
|
||||
}
|
||||
if len(docs) == 0 {
|
||||
return true, nil
|
||||
}
|
||||
ids := make([]int64, 0, len(docs))
|
||||
locationKeys := make([]string, 0, len(docs)*2)
|
||||
for id, doc := range docs {
|
||||
ids = append(ids, id)
|
||||
if doc.Path != "" {
|
||||
locationKeys = append(locationKeys, fmt.Sprintf("doc:%d", id))
|
||||
}
|
||||
for _, thumb := range doc.Thumbs {
|
||||
if thumb.Path != "" && thumb.Type != "" {
|
||||
locationKeys = append(locationKeys, fmt.Sprintf("doc:%d:%s", id, thumb.Type))
|
||||
}
|
||||
}
|
||||
}
|
||||
stored, err := s.media.GetDocuments(ctx, ids)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
if len(stored) < len(docs) {
|
||||
return false, nil
|
||||
}
|
||||
for _, doc := range stored {
|
||||
want, ok := docs[doc.ID]
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
if doc.DCID != s.dc || doc.MimeType != want.MimeType || doc.Size != want.Size {
|
||||
return false, nil
|
||||
}
|
||||
delete(docs, doc.ID)
|
||||
}
|
||||
if len(docs) > 0 {
|
||||
return false, nil
|
||||
}
|
||||
if len(locationKeys) == 0 {
|
||||
return true, nil
|
||||
}
|
||||
blobs, err := s.media.GetFileBlobs(ctx, locationKeys)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
for _, key := range locationKeys {
|
||||
if _, ok := blobs[key]; !ok {
|
||||
return false, nil
|
||||
}
|
||||
}
|
||||
return true, nil
|
||||
}
|
||||
|
||||
func (s *Service) seedAppearanceDocument(ctx context.Context, in appearance.Document) (domain.Document, int, error) {
|
||||
if in.ID == 0 {
|
||||
return domain.Document{}, 0, nil
|
||||
}
|
||||
doc := domain.Document{
|
||||
ID: in.ID,
|
||||
AccessHash: in.AccessHash,
|
||||
Date: in.Date,
|
||||
MimeType: in.MimeType,
|
||||
Size: in.Size,
|
||||
DCID: s.dc,
|
||||
Attributes: appearanceDocumentAttributes(in.Attributes),
|
||||
Thumbs: appearanceDocumentThumbs(in.Thumbs),
|
||||
}
|
||||
blobs := 0
|
||||
if in.Path != "" {
|
||||
data, sum, err := readAppearanceSeedBlob(in.Path, in.SHA256)
|
||||
if err != nil {
|
||||
return domain.Document{}, blobs, err
|
||||
}
|
||||
objectKey, err := s.blobs.Put(ctx, data)
|
||||
if err != nil {
|
||||
return domain.Document{}, blobs, err
|
||||
}
|
||||
if err := s.media.PutFileBlob(ctx, domain.FileBlob{
|
||||
LocationKey: fmt.Sprintf("doc:%d", in.ID),
|
||||
Backend: domain.MediaBackend(s.blobs.Name()),
|
||||
ObjectKey: objectKey,
|
||||
Size: int64(len(data)),
|
||||
SHA256: sum,
|
||||
MimeType: in.MimeType,
|
||||
}); err != nil {
|
||||
return domain.Document{}, blobs, err
|
||||
}
|
||||
s.prewarmSmallBlob(objectKey, data)
|
||||
blobs++
|
||||
}
|
||||
for _, thumb := range in.Thumbs {
|
||||
if thumb.Path == "" || thumb.Type == "" {
|
||||
continue
|
||||
}
|
||||
data, sum, err := readAppearanceSeedBlob(thumb.Path, thumb.SHA256)
|
||||
if err != nil {
|
||||
return domain.Document{}, blobs, err
|
||||
}
|
||||
objectKey, err := s.blobs.Put(ctx, data)
|
||||
if err != nil {
|
||||
return domain.Document{}, blobs, err
|
||||
}
|
||||
if err := s.media.PutFileBlob(ctx, domain.FileBlob{
|
||||
LocationKey: fmt.Sprintf("doc:%d:%s", in.ID, thumb.Type),
|
||||
Backend: domain.MediaBackend(s.blobs.Name()),
|
||||
ObjectKey: objectKey,
|
||||
Size: int64(len(data)),
|
||||
SHA256: sum,
|
||||
MimeType: seedThumbMimeType(data),
|
||||
}); err != nil {
|
||||
return domain.Document{}, blobs, err
|
||||
}
|
||||
s.prewarmSmallBlob(objectKey, data)
|
||||
blobs++
|
||||
}
|
||||
if err := s.media.PutDocument(ctx, doc); err != nil {
|
||||
return domain.Document{}, blobs, err
|
||||
}
|
||||
return doc, blobs, nil
|
||||
}
|
||||
|
||||
func readAppearanceSeedBlob(path, wantSHA string) ([]byte, []byte, error) {
|
||||
data, err := appearance.FS.ReadFile(path)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
sum := sha256.Sum256(data)
|
||||
got := hex.EncodeToString(sum[:])
|
||||
if wantSHA != "" && got != wantSHA {
|
||||
return nil, nil, fmt.Errorf("%s sha256 = %s, want %s", path, got, wantSHA)
|
||||
}
|
||||
return data, append([]byte(nil), sum[:]...), nil
|
||||
}
|
||||
|
||||
func appearanceDocumentAttributes(in []appearance.DocumentAttribute) []domain.DocumentAttribute {
|
||||
out := make([]domain.DocumentAttribute, 0, len(in))
|
||||
for _, attr := range in {
|
||||
switch attr.Kind {
|
||||
case "image_size":
|
||||
out = append(out, domain.DocumentAttribute{
|
||||
Kind: domain.DocAttrImageSize,
|
||||
W: attr.W,
|
||||
H: attr.H,
|
||||
})
|
||||
case "filename":
|
||||
if attr.FileName != "" {
|
||||
out = append(out, domain.DocumentAttribute{
|
||||
Kind: domain.DocAttrFilename,
|
||||
FileName: attr.FileName,
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func appearanceDocumentThumbs(in []appearance.PhotoSize) []domain.PhotoSize {
|
||||
out := make([]domain.PhotoSize, 0, len(in))
|
||||
for _, thumb := range in {
|
||||
if thumb.Type == "" {
|
||||
continue
|
||||
}
|
||||
switch thumb.Kind {
|
||||
case "size":
|
||||
out = append(out, domain.PhotoSize{
|
||||
Kind: domain.PhotoSizeKindDefault,
|
||||
Type: thumb.Type,
|
||||
W: thumb.W,
|
||||
H: thumb.H,
|
||||
Size: thumb.Size,
|
||||
})
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
98
internal/app/files/appearance_seed_test.go
Normal file
98
internal/app/files/appearance_seed_test.go
Normal file
|
|
@ -0,0 +1,98 @@
|
|||
package files
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"testing"
|
||||
|
||||
"telesrv/internal/seed/appearance"
|
||||
)
|
||||
|
||||
func TestSeedAppearanceImportsDefaultWallpaperDocuments(t *testing.T) {
|
||||
media := newFakeMediaStore()
|
||||
blobs, err := NewLocalFS(t.TempDir())
|
||||
if err != nil {
|
||||
t.Fatalf("NewLocalFS: %v", err)
|
||||
}
|
||||
svc := NewService(media, blobs, 2)
|
||||
|
||||
stats, err := svc.SeedAppearance(context.Background())
|
||||
if err != nil {
|
||||
t.Fatalf("SeedAppearance: %v", err)
|
||||
}
|
||||
if stats.Skipped || stats.Wallpapers == 0 || stats.Documents == 0 || stats.Blobs < stats.Documents {
|
||||
t.Fatalf("SeedAppearance stats = %+v, want non-empty wallpapers/documents with >=1 blob each", stats)
|
||||
}
|
||||
|
||||
var first appearance.Wallpaper
|
||||
for _, w := range appearance.Default().Wallpapers {
|
||||
if w.Document.ID != 0 {
|
||||
first = w
|
||||
break
|
||||
}
|
||||
}
|
||||
if first.Document.ID == 0 {
|
||||
t.Fatalf("no wallpaper with a document in catalog")
|
||||
}
|
||||
doc, ok, err := media.GetDocument(context.Background(), first.Document.ID)
|
||||
if err != nil || !ok {
|
||||
t.Fatalf("GetDocument(%d) = ok %v err %v", first.Document.ID, ok, err)
|
||||
}
|
||||
if doc.DCID != 2 || doc.MimeType != first.Document.MimeType || doc.Size != first.Document.Size {
|
||||
t.Fatalf("document = dc %d mime %q size %d, want dc 2 mime %q size %d",
|
||||
doc.DCID, doc.MimeType, doc.Size, first.Document.MimeType, first.Document.Size)
|
||||
}
|
||||
if len(doc.Thumbs) == 0 || doc.Thumbs[0].Type != "m" {
|
||||
t.Fatalf("document thumbs = %+v, want m thumbnail", doc.Thumbs)
|
||||
}
|
||||
if _, ok, err := media.GetFileBlob(context.Background(), fmt.Sprintf("doc:%d", first.Document.ID)); err != nil || !ok {
|
||||
t.Fatalf("main blob ok=%v err=%v, want present", ok, err)
|
||||
}
|
||||
if _, ok, err := media.GetFileBlob(context.Background(), fmt.Sprintf("doc:%d:m", first.Document.ID)); err != nil || !ok {
|
||||
t.Fatalf("thumb blob ok=%v err=%v, want present", ok, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSeedAppearanceSkipsUnchangedCatalogAndRepairsMissingBlob(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
media := newFakeMediaStore()
|
||||
blobs, err := NewLocalFS(t.TempDir())
|
||||
if err != nil {
|
||||
t.Fatalf("NewLocalFS: %v", err)
|
||||
}
|
||||
svc := NewService(media, blobs, 2)
|
||||
|
||||
first, err := svc.SeedAppearance(ctx)
|
||||
if err != nil {
|
||||
t.Fatalf("first SeedAppearance: %v", err)
|
||||
}
|
||||
if first.Skipped || first.Documents == 0 || first.Blobs == 0 {
|
||||
t.Fatalf("first stats = %+v, want import", first)
|
||||
}
|
||||
second, err := svc.SeedAppearance(ctx)
|
||||
if err != nil {
|
||||
t.Fatalf("second SeedAppearance: %v", err)
|
||||
}
|
||||
if !second.Skipped || second.Documents != 0 || second.Blobs != 0 {
|
||||
t.Fatalf("second stats = %+v, want unchanged catalog skip", second)
|
||||
}
|
||||
|
||||
var firstDocID int64
|
||||
for _, w := range appearance.Default().Wallpapers {
|
||||
if w.Document.ID != 0 && w.Document.Path != "" {
|
||||
firstDocID = w.Document.ID
|
||||
break
|
||||
}
|
||||
}
|
||||
if firstDocID == 0 {
|
||||
t.Fatal("no wallpaper document found")
|
||||
}
|
||||
delete(media.blobs, fmt.Sprintf("doc:%d", firstDocID))
|
||||
repaired, err := svc.SeedAppearance(ctx)
|
||||
if err != nil {
|
||||
t.Fatalf("repair SeedAppearance: %v", err)
|
||||
}
|
||||
if repaired.Skipped || repaired.Documents == 0 || repaired.Blobs == 0 {
|
||||
t.Fatalf("repair stats = %+v, want missing blob to force reimport", repaired)
|
||||
}
|
||||
}
|
||||
|
|
@ -2,11 +2,74 @@ package files
|
|||
|
||||
import (
|
||||
"container/list"
|
||||
"strconv"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"telesrv/internal/domain"
|
||||
)
|
||||
|
||||
// stickerSetNegativeCache 是「按 ref 查不到的贴纸集」的短 TTL 负缓存。未 seed 的 short_name
|
||||
// 集合会被客户端反复 getStickerSet(每次都打一发 PG GetStickerSetByShortName),这里缓存
|
||||
// not-found 结果,TTL 内直接短路、不再查库。TTL 短(自愈):运营/运行时新增的集合最多 TTL 后
|
||||
// 才被解析,避免负缓存长期遮住真实存在的集合。
|
||||
type stickerSetNegativeCache struct {
|
||||
mu sync.Mutex
|
||||
ttl time.Duration
|
||||
entries map[string]time.Time
|
||||
}
|
||||
|
||||
const stickerSetNegativeCacheMaxEntries = 100000
|
||||
|
||||
func newStickerSetNegativeCache(ttl time.Duration) *stickerSetNegativeCache {
|
||||
return &stickerSetNegativeCache{ttl: ttl, entries: map[string]time.Time{}}
|
||||
}
|
||||
|
||||
func stickerSetRefKey(ref domain.StickerSetRef) string {
|
||||
switch ref.Kind {
|
||||
case domain.StickerSetRefByID:
|
||||
return "id:" + strconv.FormatInt(ref.ID, 10)
|
||||
case domain.StickerSetRefByShortName:
|
||||
return "short:" + ref.ShortName
|
||||
case domain.StickerSetRefBySystem:
|
||||
return "sys:" + ref.SystemKey
|
||||
default:
|
||||
return ""
|
||||
}
|
||||
}
|
||||
|
||||
func (c *stickerSetNegativeCache) has(ref domain.StickerSetRef) bool {
|
||||
key := stickerSetRefKey(ref)
|
||||
if key == "" {
|
||||
return false
|
||||
}
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
exp, ok := c.entries[key]
|
||||
if !ok {
|
||||
return false
|
||||
}
|
||||
if time.Now().After(exp) {
|
||||
delete(c.entries, key)
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
func (c *stickerSetNegativeCache) put(ref domain.StickerSetRef) {
|
||||
key := stickerSetRefKey(ref)
|
||||
if key == "" {
|
||||
return
|
||||
}
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
// 简单上限防无界增长:超限整表清空(短 TTL 下冷启代价可忽略)。
|
||||
if len(c.entries) >= stickerSetNegativeCacheMaxEntries {
|
||||
c.entries = make(map[string]time.Time, 1024)
|
||||
}
|
||||
c.entries[key] = time.Now().Add(c.ttl)
|
||||
}
|
||||
|
||||
// blobMetaCache 是 location_key → FileBlob 元数据的进程内 LRU,用于消除 upload.getFile
|
||||
// 每个 chunk 一次 GetFileBlob 的 PG 往返(一个文件按 ≤512KB/1MB 分多次 getFile,热门贴纸/
|
||||
// reaction/头像更被大量用户重复拉)。
|
||||
|
|
|
|||
|
|
@ -1,13 +1,20 @@
|
|||
package files
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"fmt"
|
||||
"io"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"telesrv/internal/domain"
|
||||
)
|
||||
|
||||
// BlobBackend 是 blob 字节内容的存储后端。第一阶段只有本地磁盘实现。
|
||||
|
|
@ -15,15 +22,42 @@ import (
|
|||
type BlobBackend interface {
|
||||
Name() string
|
||||
Put(ctx context.Context, data []byte) (objectKey string, err error)
|
||||
PutReader(ctx context.Context, r io.Reader) (objectKey string, size int64, sha256 []byte, err error)
|
||||
Get(ctx context.Context, objectKey string) ([]byte, error)
|
||||
// GetRange 只读 [offset, offset+limit) 段并返回该段字节与文件总大小(limit<=0 读到末尾),
|
||||
// 避免大文件每个 chunk 都整文件读入内存(getFile 按 chunk 多次请求 ⇒ 否则 O(N²) 放大)。
|
||||
GetRange(ctx context.Context, objectKey string, offset, limit int64) (data []byte, total int64, err error)
|
||||
}
|
||||
|
||||
// UploadPartBackend 保存 upload.saveFilePart/saveBigFilePart 的临时分片字节。
|
||||
// 与正式 blob 不同,上传分片 key 唯一且可删除,成功组装/覆盖重传/GC 后必须清理。
|
||||
type UploadPartBackend interface {
|
||||
PutUploadPart(ctx context.Context, ownerUserID, fileID int64, part int, data []byte) (uploadPartObject, error)
|
||||
GetUploadPart(ctx context.Context, objectKey string) ([]byte, error)
|
||||
OpenUploadPart(ctx context.Context, objectKey string) (io.ReadCloser, error)
|
||||
DeleteUploadPart(ctx context.Context, objectKey string) error
|
||||
DeleteExpiredUploadParts(ctx context.Context, before time.Time, limit int) (int64, error)
|
||||
}
|
||||
|
||||
type uploadPartObject struct {
|
||||
Backend domain.MediaBackend
|
||||
ObjectKey string
|
||||
Size int64
|
||||
SHA256 []byte
|
||||
}
|
||||
|
||||
// LocalFS 把 blob 字节存到本地磁盘根目录下,路径按内容 hash 两级 fanout。
|
||||
type LocalFS struct {
|
||||
root string
|
||||
|
||||
mu sync.Mutex
|
||||
openBlobFiles map[string]*sharedBlobFile
|
||||
}
|
||||
|
||||
type sharedBlobFile struct {
|
||||
key string
|
||||
file *os.File
|
||||
refs int
|
||||
}
|
||||
|
||||
// NewLocalFS 创建本地磁盘 blob backend,确保根目录存在。
|
||||
|
|
@ -34,7 +68,7 @@ func NewLocalFS(root string) (*LocalFS, error) {
|
|||
if err := os.MkdirAll(root, 0o755); err != nil {
|
||||
return nil, fmt.Errorf("create blob root %q: %w", root, err)
|
||||
}
|
||||
return &LocalFS{root: root}, nil
|
||||
return &LocalFS{root: root, openBlobFiles: make(map[string]*sharedBlobFile)}, nil
|
||||
}
|
||||
|
||||
// Name 返回后端标识,与 file_blobs.backend 一致。
|
||||
|
|
@ -48,24 +82,90 @@ func (l *LocalFS) pathFor(objectKey string) string {
|
|||
}
|
||||
|
||||
// Put 写入内容并返回 sha256 hex 作为 objectKey;同内容已存在则跳过写入(去重)。
|
||||
func (l *LocalFS) Put(_ context.Context, data []byte) (string, error) {
|
||||
sum := sha256.Sum256(data)
|
||||
key := hex.EncodeToString(sum[:])
|
||||
func (l *LocalFS) Put(ctx context.Context, data []byte) (string, error) {
|
||||
key, _, _, err := l.PutReader(ctx, bytes.NewReader(data))
|
||||
return key, err
|
||||
}
|
||||
|
||||
// PutReader 流式写入内容,边复制边计算 sha256,避免上层为大视频先拼出完整 []byte。
|
||||
func (l *LocalFS) PutReader(ctx context.Context, r io.Reader) (string, int64, []byte, error) {
|
||||
tmpDir := filepath.Join(l.root, "_tmp")
|
||||
if err := os.MkdirAll(tmpDir, 0o755); err != nil {
|
||||
return "", 0, nil, fmt.Errorf("create blob tmp dir: %w", err)
|
||||
}
|
||||
tmp, err := os.CreateTemp(tmpDir, "blob-*.tmp")
|
||||
if err != nil {
|
||||
return "", 0, nil, fmt.Errorf("create blob tmp file: %w", err)
|
||||
}
|
||||
tmpPath := tmp.Name()
|
||||
committed := false
|
||||
defer func() {
|
||||
if !committed {
|
||||
_ = os.Remove(tmpPath)
|
||||
}
|
||||
}()
|
||||
|
||||
h := sha256.New()
|
||||
size, err := copyWithContext(ctx, io.MultiWriter(tmp, h), r)
|
||||
closeErr := tmp.Close()
|
||||
if err != nil {
|
||||
return "", 0, nil, fmt.Errorf("write blob stream: %w", err)
|
||||
}
|
||||
if closeErr != nil {
|
||||
return "", 0, nil, fmt.Errorf("close blob stream: %w", closeErr)
|
||||
}
|
||||
sum := h.Sum(nil)
|
||||
key := hex.EncodeToString(sum)
|
||||
path := l.pathFor(key)
|
||||
if _, err := os.Stat(path); err == nil {
|
||||
return key, nil
|
||||
committed = true
|
||||
_ = os.Remove(tmpPath)
|
||||
return key, size, append([]byte(nil), sum...), nil
|
||||
} else if err != nil && !os.IsNotExist(err) {
|
||||
return "", 0, nil, fmt.Errorf("stat blob: %w", err)
|
||||
}
|
||||
if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil {
|
||||
return "", fmt.Errorf("create blob dir: %w", err)
|
||||
return "", 0, nil, fmt.Errorf("create blob dir: %w", err)
|
||||
}
|
||||
tmp := path + ".tmp"
|
||||
if err := os.WriteFile(tmp, data, 0o644); err != nil {
|
||||
return "", fmt.Errorf("write blob: %w", err)
|
||||
if err := os.Rename(tmpPath, path); err != nil {
|
||||
if _, statErr := os.Stat(path); statErr == nil {
|
||||
committed = true
|
||||
_ = os.Remove(tmpPath)
|
||||
return key, size, append([]byte(nil), sum...), nil
|
||||
}
|
||||
return "", 0, nil, fmt.Errorf("commit blob: %w", err)
|
||||
}
|
||||
if err := os.Rename(tmp, path); err != nil {
|
||||
return "", fmt.Errorf("commit blob: %w", err)
|
||||
committed = true
|
||||
return key, size, append([]byte(nil), sum...), nil
|
||||
}
|
||||
|
||||
func copyWithContext(ctx context.Context, dst io.Writer, src io.Reader) (int64, error) {
|
||||
buf := make([]byte, 256<<10)
|
||||
var written int64
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return written, ctx.Err()
|
||||
default:
|
||||
}
|
||||
n, readErr := src.Read(buf)
|
||||
if n > 0 {
|
||||
w, writeErr := dst.Write(buf[:n])
|
||||
written += int64(w)
|
||||
if writeErr != nil {
|
||||
return written, writeErr
|
||||
}
|
||||
if w != n {
|
||||
return written, io.ErrShortWrite
|
||||
}
|
||||
}
|
||||
if readErr == io.EOF {
|
||||
return written, nil
|
||||
}
|
||||
if readErr != nil {
|
||||
return written, readErr
|
||||
}
|
||||
}
|
||||
return key, nil
|
||||
}
|
||||
|
||||
// Get 读取 objectKey 对应的全部字节。
|
||||
|
|
@ -76,11 +176,12 @@ func (l *LocalFS) Get(_ context.Context, objectKey string) ([]byte, error) {
|
|||
// GetRange 用 ReadAt 只读 [offset, offset+limit) 段,total 取自文件大小;
|
||||
// n 受 total 约束,故即便客户端传超大 limit 也只分配文件实际大小,不会按客户端巨值分配。
|
||||
func (l *LocalFS) GetRange(_ context.Context, objectKey string, offset, limit int64) ([]byte, int64, error) {
|
||||
f, err := os.Open(l.pathFor(objectKey))
|
||||
blobFile, err := l.openBlobFile(objectKey)
|
||||
if err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
defer f.Close()
|
||||
defer l.releaseBlobFile(blobFile)
|
||||
f := blobFile.file
|
||||
info, err := f.Stat()
|
||||
if err != nil {
|
||||
return nil, 0, err
|
||||
|
|
@ -103,3 +204,157 @@ func (l *LocalFS) GetRange(_ context.Context, objectKey string, offset, limit in
|
|||
}
|
||||
return buf[:read], total, nil
|
||||
}
|
||||
|
||||
func (l *LocalFS) openBlobFile(objectKey string) (*sharedBlobFile, error) {
|
||||
l.mu.Lock()
|
||||
defer l.mu.Unlock()
|
||||
if f, ok := l.openBlobFiles[objectKey]; ok {
|
||||
f.refs++
|
||||
return f, nil
|
||||
}
|
||||
f, err := os.Open(l.pathFor(objectKey))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
blobFile := &sharedBlobFile{
|
||||
key: objectKey,
|
||||
file: f,
|
||||
refs: 1,
|
||||
}
|
||||
l.openBlobFiles[objectKey] = blobFile
|
||||
return blobFile, nil
|
||||
}
|
||||
|
||||
func (l *LocalFS) releaseBlobFile(blobFile *sharedBlobFile) {
|
||||
if blobFile == nil {
|
||||
return
|
||||
}
|
||||
l.mu.Lock()
|
||||
defer l.mu.Unlock()
|
||||
if blobFile.refs > 0 {
|
||||
blobFile.refs--
|
||||
}
|
||||
if blobFile.refs != 0 {
|
||||
return
|
||||
}
|
||||
if current := l.openBlobFiles[blobFile.key]; current == blobFile {
|
||||
delete(l.openBlobFiles, blobFile.key)
|
||||
}
|
||||
_ = blobFile.file.Close()
|
||||
}
|
||||
|
||||
func (l *LocalFS) PutUploadPart(_ context.Context, ownerUserID, fileID int64, part int, data []byte) (uploadPartObject, error) {
|
||||
sum := sha256.Sum256(data)
|
||||
var nonce [16]byte
|
||||
if _, err := rand.Read(nonce[:]); err != nil {
|
||||
return uploadPartObject{}, fmt.Errorf("generate upload part key: %w", err)
|
||||
}
|
||||
key := filepath.ToSlash(filepath.Join(
|
||||
"upload_parts",
|
||||
fmt.Sprintf("%d", ownerUserID),
|
||||
fmt.Sprintf("%d", fileID),
|
||||
fmt.Sprintf("%06d-%s.part", part, hex.EncodeToString(nonce[:])),
|
||||
))
|
||||
path, err := l.uploadPartPath(key)
|
||||
if err != nil {
|
||||
return uploadPartObject{}, err
|
||||
}
|
||||
if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil {
|
||||
return uploadPartObject{}, fmt.Errorf("create upload part dir: %w", err)
|
||||
}
|
||||
tmp := path + ".tmp"
|
||||
if err := os.WriteFile(tmp, data, 0o644); err != nil {
|
||||
return uploadPartObject{}, fmt.Errorf("write upload part: %w", err)
|
||||
}
|
||||
if err := os.Rename(tmp, path); err != nil {
|
||||
_ = os.Remove(tmp)
|
||||
return uploadPartObject{}, fmt.Errorf("commit upload part: %w", err)
|
||||
}
|
||||
return uploadPartObject{
|
||||
Backend: domain.MediaBackend(l.Name()),
|
||||
ObjectKey: key,
|
||||
Size: int64(len(data)),
|
||||
SHA256: append([]byte(nil), sum[:]...),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (l *LocalFS) GetUploadPart(_ context.Context, objectKey string) ([]byte, error) {
|
||||
path, err := l.uploadPartPath(objectKey)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return os.ReadFile(path)
|
||||
}
|
||||
|
||||
func (l *LocalFS) OpenUploadPart(_ context.Context, objectKey string) (io.ReadCloser, error) {
|
||||
path, err := l.uploadPartPath(objectKey)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return os.Open(path)
|
||||
}
|
||||
|
||||
func (l *LocalFS) DeleteUploadPart(_ context.Context, objectKey string) error {
|
||||
path, err := l.uploadPartPath(objectKey)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := os.Remove(path); err != nil && !os.IsNotExist(err) {
|
||||
return fmt.Errorf("delete upload part: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (l *LocalFS) DeleteExpiredUploadParts(ctx context.Context, before time.Time, limit int) (int64, error) {
|
||||
if limit <= 0 {
|
||||
return 0, nil
|
||||
}
|
||||
root := filepath.Join(l.root, "upload_parts")
|
||||
if _, err := os.Stat(root); os.IsNotExist(err) {
|
||||
return 0, nil
|
||||
} else if err != nil {
|
||||
return 0, fmt.Errorf("stat upload parts root: %w", err)
|
||||
}
|
||||
var deleted int64
|
||||
err := filepath.WalkDir(root, func(path string, d os.DirEntry, err error) error {
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if deleted >= int64(limit) {
|
||||
return filepath.SkipAll
|
||||
}
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return ctx.Err()
|
||||
default:
|
||||
}
|
||||
if d.IsDir() {
|
||||
return nil
|
||||
}
|
||||
info, err := d.Info()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if !info.ModTime().Before(before) {
|
||||
return nil
|
||||
}
|
||||
if err := os.Remove(path); err != nil && !os.IsNotExist(err) {
|
||||
return err
|
||||
}
|
||||
deleted++
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
return deleted, fmt.Errorf("delete expired upload part objects: %w", err)
|
||||
}
|
||||
return deleted, nil
|
||||
}
|
||||
|
||||
func (l *LocalFS) uploadPartPath(objectKey string) (string, error) {
|
||||
clean := filepath.Clean(filepath.FromSlash(objectKey))
|
||||
prefix := "upload_parts" + string(os.PathSeparator)
|
||||
if clean == "." || clean == ".." || filepath.IsAbs(clean) || strings.HasPrefix(clean, ".."+string(os.PathSeparator)) || !strings.HasPrefix(clean, prefix) {
|
||||
return "", fmt.Errorf("invalid upload part object key %q", objectKey)
|
||||
}
|
||||
return filepath.Join(l.root, clean), nil
|
||||
}
|
||||
|
|
|
|||
|
|
@ -3,7 +3,11 @@ package files
|
|||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"os"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestLocalFSPutGetRoundTrip(t *testing.T) {
|
||||
|
|
@ -40,6 +44,68 @@ func TestLocalFSPutGetRoundTrip(t *testing.T) {
|
|||
}
|
||||
}
|
||||
|
||||
func TestLocalFSPutReaderRoundTrip(t *testing.T) {
|
||||
fs, err := NewLocalFS(t.TempDir())
|
||||
if err != nil {
|
||||
t.Fatalf("new local fs: %v", err)
|
||||
}
|
||||
ctx := context.Background()
|
||||
data := strings.Repeat("streamed-", 1024)
|
||||
|
||||
key, size, sum, err := fs.PutReader(ctx, strings.NewReader(data))
|
||||
if err != nil {
|
||||
t.Fatalf("put reader: %v", err)
|
||||
}
|
||||
if key == "" || size != int64(len(data)) || len(sum) != 32 {
|
||||
t.Fatalf("stream metadata key=%q size=%d sha=%d", key, size, len(sum))
|
||||
}
|
||||
got, err := fs.Get(ctx, key)
|
||||
if err != nil {
|
||||
t.Fatalf("get: %v", err)
|
||||
}
|
||||
if string(got) != data {
|
||||
t.Fatalf("roundtrip mismatch")
|
||||
}
|
||||
}
|
||||
|
||||
func TestLocalFSDeleteExpiredUploadParts(t *testing.T) {
|
||||
fs, err := NewLocalFS(t.TempDir())
|
||||
if err != nil {
|
||||
t.Fatalf("new local fs: %v", err)
|
||||
}
|
||||
ctx := context.Background()
|
||||
oldPart, err := fs.PutUploadPart(ctx, 10, 100, 0, []byte("old"))
|
||||
if err != nil {
|
||||
t.Fatalf("put old upload part: %v", err)
|
||||
}
|
||||
freshPart, err := fs.PutUploadPart(ctx, 10, 100, 1, []byte("fresh"))
|
||||
if err != nil {
|
||||
t.Fatalf("put fresh upload part: %v", err)
|
||||
}
|
||||
oldPath, err := fs.uploadPartPath(oldPart.ObjectKey)
|
||||
if err != nil {
|
||||
t.Fatalf("old upload part path: %v", err)
|
||||
}
|
||||
oldTime := time.Now().Add(-48 * time.Hour)
|
||||
if err := os.Chtimes(oldPath, oldTime, oldTime); err != nil {
|
||||
t.Fatalf("age old upload part: %v", err)
|
||||
}
|
||||
|
||||
deleted, err := fs.DeleteExpiredUploadParts(ctx, time.Now().Add(-24*time.Hour), 10)
|
||||
if err != nil {
|
||||
t.Fatalf("delete expired upload parts: %v", err)
|
||||
}
|
||||
if deleted != 1 {
|
||||
t.Fatalf("deleted = %d, want 1", deleted)
|
||||
}
|
||||
if _, err := fs.GetUploadPart(ctx, oldPart.ObjectKey); err == nil {
|
||||
t.Fatalf("old upload part still exists")
|
||||
}
|
||||
if data, err := fs.GetUploadPart(ctx, freshPart.ObjectKey); err != nil || string(data) != "fresh" {
|
||||
t.Fatalf("fresh upload part = %q err=%v", data, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLocalFSDistinctContent(t *testing.T) {
|
||||
fs, err := NewLocalFS(t.TempDir())
|
||||
if err != nil {
|
||||
|
|
@ -90,3 +156,143 @@ func TestLocalFSGetRange(t *testing.T) {
|
|||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestLocalFSReusesOpenBlobFileWhileActive(t *testing.T) {
|
||||
fs, err := NewLocalFS(t.TempDir())
|
||||
if err != nil {
|
||||
t.Fatalf("new local fs: %v", err)
|
||||
}
|
||||
ctx := context.Background()
|
||||
key, err := fs.Put(ctx, []byte("0123456789"))
|
||||
if err != nil {
|
||||
t.Fatalf("put: %v", err)
|
||||
}
|
||||
|
||||
first, err := fs.openBlobFile(key)
|
||||
if err != nil {
|
||||
t.Fatalf("open first: %v", err)
|
||||
}
|
||||
|
||||
second, err := fs.openBlobFile(key)
|
||||
if err != nil {
|
||||
t.Fatalf("open second: %v", err)
|
||||
}
|
||||
firstReleased, secondReleased := false, false
|
||||
t.Cleanup(func() {
|
||||
if !firstReleased {
|
||||
fs.releaseBlobFile(first)
|
||||
}
|
||||
if !secondReleased {
|
||||
fs.releaseBlobFile(second)
|
||||
}
|
||||
})
|
||||
if first != second {
|
||||
t.Fatal("same active blob should reuse one open file handle")
|
||||
}
|
||||
fs.mu.Lock()
|
||||
open, refs := len(fs.openBlobFiles), first.refs
|
||||
fs.mu.Unlock()
|
||||
if open != 1 || refs != 2 {
|
||||
t.Fatalf("active files=%d refs=%d, want 1/2", open, refs)
|
||||
}
|
||||
|
||||
fs.releaseBlobFile(first)
|
||||
firstReleased = true
|
||||
fs.mu.Lock()
|
||||
open, refs = len(fs.openBlobFiles), second.refs
|
||||
fs.mu.Unlock()
|
||||
if open != 1 || refs != 1 {
|
||||
t.Fatalf("after first release active files=%d refs=%d, want 1/1", open, refs)
|
||||
}
|
||||
|
||||
var buf [4]byte
|
||||
if n, err := second.file.ReadAt(buf[:], 3); err != nil || n != len(buf) || string(buf[:]) != "3456" {
|
||||
t.Fatalf("shared file ReadAt n=%d err=%v bytes=%q, want 3456", n, err, buf[:])
|
||||
}
|
||||
|
||||
fs.releaseBlobFile(second)
|
||||
secondReleased = true
|
||||
fs.mu.Lock()
|
||||
open = len(fs.openBlobFiles)
|
||||
fs.mu.Unlock()
|
||||
if open != 0 {
|
||||
t.Fatalf("after final release active files=%d, want 0", open)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLocalFSReusesOpenBlobFileUnderConcurrentOpen(t *testing.T) {
|
||||
fs, err := NewLocalFS(t.TempDir())
|
||||
if err != nil {
|
||||
t.Fatalf("new local fs: %v", err)
|
||||
}
|
||||
ctx := context.Background()
|
||||
key, err := fs.Put(ctx, []byte(strings.Repeat("x", 1024)))
|
||||
if err != nil {
|
||||
t.Fatalf("put: %v", err)
|
||||
}
|
||||
|
||||
const readers = 32
|
||||
start := make(chan struct{})
|
||||
files := make([]*sharedBlobFile, readers)
|
||||
errs := make([]error, readers)
|
||||
var wg sync.WaitGroup
|
||||
for i := range files {
|
||||
wg.Add(1)
|
||||
go func(i int) {
|
||||
defer wg.Done()
|
||||
<-start
|
||||
files[i], errs[i] = fs.openBlobFile(key)
|
||||
}(i)
|
||||
}
|
||||
close(start)
|
||||
wg.Wait()
|
||||
|
||||
released := false
|
||||
t.Cleanup(func() {
|
||||
if released {
|
||||
return
|
||||
}
|
||||
for _, f := range files {
|
||||
if f != nil {
|
||||
fs.releaseBlobFile(f)
|
||||
}
|
||||
}
|
||||
})
|
||||
for i, err := range errs {
|
||||
if err != nil {
|
||||
t.Fatalf("open %d: %v", i, err)
|
||||
}
|
||||
}
|
||||
first := files[0]
|
||||
if first == nil {
|
||||
t.Fatal("first open returned nil")
|
||||
}
|
||||
for i, f := range files {
|
||||
if f != first {
|
||||
t.Fatalf("file %d = %p, want shared %p", i, f, first)
|
||||
}
|
||||
}
|
||||
fs.mu.Lock()
|
||||
open, refs := len(fs.openBlobFiles), first.refs
|
||||
fs.mu.Unlock()
|
||||
if open != 1 || refs != readers {
|
||||
t.Fatalf("active files=%d refs=%d, want 1/%d", open, refs, readers)
|
||||
}
|
||||
|
||||
wg = sync.WaitGroup{}
|
||||
for _, f := range files {
|
||||
wg.Add(1)
|
||||
go func(f *sharedBlobFile) {
|
||||
defer wg.Done()
|
||||
fs.releaseBlobFile(f)
|
||||
}(f)
|
||||
}
|
||||
wg.Wait()
|
||||
released = true
|
||||
fs.mu.Lock()
|
||||
open = len(fs.openBlobFiles)
|
||||
fs.mu.Unlock()
|
||||
if open != 0 {
|
||||
t.Fatalf("after concurrent release active files=%d, want 0", open)
|
||||
}
|
||||
}
|
||||
|
|
|
|||
148
internal/app/files/default_statuses.go
Normal file
148
internal/app/files/default_statuses.go
Normal file
|
|
@ -0,0 +1,148 @@
|
|||
package files
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"telesrv/internal/domain"
|
||||
)
|
||||
|
||||
// 合成集的固定 ID/AccessHash:种子导出的真实集 ID 在 1.2e15 量级,取明显隔离的
|
||||
// 常量段避免撞键;幂等性由 system_key 查询保证,常量仅需稳定。
|
||||
const (
|
||||
defaultEmojiStatusSetID int64 = 7_777_000_000_000_001
|
||||
defaultEmojiStatusSetAccessHash int64 = 7_777_000_000_000_002
|
||||
)
|
||||
|
||||
// defaultEmojiStatusEmoticons 是默认状态的精选 emoji(对齐官方默认状态选盘的
|
||||
// 常见项),按展示顺序排列。匹配不到的 emoticon 静默跳过(取决于 seed 内容)。
|
||||
var defaultEmojiStatusEmoticons = []string{
|
||||
"\U0001f4bc", // 💼 工作
|
||||
"\U0001f393", // 🎓 学习
|
||||
"\U0001f3e0", // 🏠 在家
|
||||
"\U0001f334", // 🌴 度假
|
||||
"\U0001f3d6", // 🏖 海滩
|
||||
"✈", // ✈️ 旅行
|
||||
"\U0001f912", // 🤒 生病
|
||||
"\U0001f634", // 😴 睡觉
|
||||
"☕", // ☕ 咖啡
|
||||
"\U0001f4bb", // 💻 编码/办公
|
||||
"\U0001f4da", // 📚 阅读
|
||||
"\U0001f3ae", // 🎮 游戏
|
||||
"\U0001f3a7", // 🎧 听歌
|
||||
"⚽", // ⚽ 运动
|
||||
"\U0001f3c6", // 🏆 获胜
|
||||
"❤", // ❤️ 爱心
|
||||
"\U0001f60e", // 😎 酷
|
||||
"\U0001f319", // 🌙 勿扰
|
||||
"⭐", // ⭐ 星标
|
||||
"\U0001f525", // 🔥 火
|
||||
"\U0001f44d", // 👍 赞
|
||||
"\U0001f389", // 🎉 庆祝
|
||||
"\U0001f914", // 🤔 思考
|
||||
"\U0001f607", // 😇 天使
|
||||
"\U0001f973", // 🥳 派对
|
||||
"\U0001f602", // 😂 大笑
|
||||
"\U0001f970", // 🥰 喜爱
|
||||
"\U0001f62d", // 😭 大哭
|
||||
"\U0001f92f", // 🤯 爆炸
|
||||
"\U0001f440", // 👀 围观
|
||||
"\U0001f4af", // 💯 满分
|
||||
"\U0001f64f", // 🙏 感谢
|
||||
"\U0001f91d", // 🤝 合作
|
||||
"✍", // ✍️ 写作
|
||||
"\U0001f697", // 🚗 通勤
|
||||
"\U0001f355", // 🍕 干饭
|
||||
"\U0001f382", // 🎂 生日
|
||||
"\U0001f338", // 🌸 春天
|
||||
"⛄", // ⛄ 冬天
|
||||
"\U0001f984", // 🦄 独角兽
|
||||
}
|
||||
|
||||
// EnsureDefaultEmojiStatusSet 幂等地合成默认 emoji status 系统集:从已 seed 的
|
||||
// animated_emoji 系统集按 emoticon 精选文档(复用文档行与 blob,不复制字节)。
|
||||
// 返回 (集内文档数, 是否本次新建)。animated_emoji 未 seed 时静默跳过。
|
||||
func (s *Service) EnsureDefaultEmojiStatusSet(ctx context.Context) (int, bool, error) {
|
||||
if existing, found, err := s.media.GetStickerSetBySystemKey(ctx, domain.StickerSetSystemKeyEmojiDefaultStatuses); err != nil {
|
||||
return 0, false, fmt.Errorf("lookup default emoji status set: %w", err)
|
||||
} else if found {
|
||||
return len(existing.DocumentIDs), false, nil
|
||||
}
|
||||
source, found, err := s.media.GetStickerSetBySystemKey(ctx, "animated_emoji")
|
||||
if err != nil {
|
||||
return 0, false, fmt.Errorf("lookup animated_emoji set: %w", err)
|
||||
}
|
||||
if !found || len(source.Packs) == 0 {
|
||||
return 0, false, nil
|
||||
}
|
||||
byEmoticon := make(map[string][]int64, len(source.Packs))
|
||||
for _, pack := range source.Packs {
|
||||
key := normalizeStatusEmoticon(pack.Emoticon)
|
||||
if key == "" || len(pack.DocumentIDs) == 0 {
|
||||
continue
|
||||
}
|
||||
byEmoticon[key] = append(byEmoticon[key], pack.DocumentIDs...)
|
||||
}
|
||||
set := domain.StickerSet{
|
||||
ID: defaultEmojiStatusSetID,
|
||||
AccessHash: defaultEmojiStatusSetAccessHash,
|
||||
ShortName: "TelesrvDefaultStatuses",
|
||||
Title: "Default Emoji Statuses",
|
||||
Kind: domain.StickerSetKindSystem,
|
||||
SystemKey: domain.StickerSetSystemKeyEmojiDefaultStatuses,
|
||||
Official: true,
|
||||
Animated: source.Animated,
|
||||
Emojis: true,
|
||||
}
|
||||
seen := make(map[int64]struct{})
|
||||
for _, emoticon := range defaultEmojiStatusEmoticons {
|
||||
ids := byEmoticon[normalizeStatusEmoticon(emoticon)]
|
||||
if len(ids) == 0 {
|
||||
continue
|
||||
}
|
||||
pack := domain.StickerPack{Emoticon: emoticon}
|
||||
for _, id := range ids {
|
||||
if _, ok := seen[id]; ok {
|
||||
continue
|
||||
}
|
||||
seen[id] = struct{}{}
|
||||
set.DocumentIDs = append(set.DocumentIDs, id)
|
||||
pack.DocumentIDs = append(pack.DocumentIDs, id)
|
||||
}
|
||||
if len(pack.DocumentIDs) > 0 {
|
||||
set.Packs = append(set.Packs, pack)
|
||||
}
|
||||
}
|
||||
if len(set.DocumentIDs) == 0 {
|
||||
return 0, false, nil
|
||||
}
|
||||
set.Count = len(set.DocumentIDs)
|
||||
set.Hash = stickerSetDocsHash(set.DocumentIDs)
|
||||
docs, err := s.media.GetDocuments(ctx, set.DocumentIDs)
|
||||
if err != nil {
|
||||
return 0, false, fmt.Errorf("load default emoji status documents: %w", err)
|
||||
}
|
||||
if err := s.media.PutStickerSet(ctx, set); err != nil {
|
||||
return 0, false, fmt.Errorf("persist default emoji status set: %w", err)
|
||||
}
|
||||
s.stickerSetCache.put(set, orderDocuments(docs, set.DocumentIDs))
|
||||
return set.Count, true, nil
|
||||
}
|
||||
|
||||
// normalizeStatusEmoticon 统一变体选择符差异("❤" vs "❤️"),匹配 seed packs 与
|
||||
// 精选清单两侧的书写形态。
|
||||
func normalizeStatusEmoticon(e string) string {
|
||||
return strings.ReplaceAll(strings.TrimSpace(e), "️", "")
|
||||
}
|
||||
|
||||
// stickerSetDocsHash 由文档 ID 列表算稳定 set hash(messages.getStickerSet 与
|
||||
// account.getDefaultEmojiStatuses 共用一份缓存判定)。
|
||||
func stickerSetDocsHash(ids []int64) int {
|
||||
var h uint64
|
||||
for _, id := range ids {
|
||||
h ^= uint64(id)
|
||||
h = h*0x4f25 + uint64(id)
|
||||
}
|
||||
return int(h & 0x7fffffff)
|
||||
}
|
||||
117
internal/app/files/default_statuses_test.go
Normal file
117
internal/app/files/default_statuses_test.go
Normal file
|
|
@ -0,0 +1,117 @@
|
|||
package files
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"telesrv/internal/domain"
|
||||
)
|
||||
|
||||
func defaultStatusesTestService(t *testing.T) (*Service, *fakeMediaStore) {
|
||||
t.Helper()
|
||||
media := newFakeMediaStore()
|
||||
return NewService(media, nil, 2), media
|
||||
}
|
||||
|
||||
func putAnimatedEmojiSet(t *testing.T, media *fakeMediaStore, packs []domain.StickerPack) {
|
||||
t.Helper()
|
||||
ctx := context.Background()
|
||||
var ids []int64
|
||||
for _, p := range packs {
|
||||
ids = append(ids, p.DocumentIDs...)
|
||||
}
|
||||
for _, id := range ids {
|
||||
if err := media.PutDocument(ctx, domain.Document{ID: id, AccessHash: id, MimeType: "application/x-tgsticker"}); err != nil {
|
||||
t.Fatalf("put document %d: %v", id, err)
|
||||
}
|
||||
}
|
||||
if err := media.PutStickerSet(ctx, domain.StickerSet{
|
||||
ID: 1,
|
||||
ShortName: "AnimatedEmojies",
|
||||
Kind: domain.StickerSetKindSystem,
|
||||
SystemKey: "animated_emoji",
|
||||
Animated: true,
|
||||
DocumentIDs: ids,
|
||||
Packs: packs,
|
||||
}); err != nil {
|
||||
t.Fatalf("put animated_emoji set: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEnsureDefaultEmojiStatusSet(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
svc, media := defaultStatusesTestService(t)
|
||||
putAnimatedEmojiSet(t, media, []domain.StickerPack{
|
||||
// seed 导出常见为裸码点(无 FE0F),精选清单两种形态都必须匹配。
|
||||
{Emoticon: "❤", DocumentIDs: []int64{101}},
|
||||
{Emoticon: "👍", DocumentIDs: []int64{102}},
|
||||
{Emoticon: "☕️", DocumentIDs: []int64{103}}, // 带 FE0F 的反向形态
|
||||
{Emoticon: "🥔", DocumentIDs: []int64{999}}, // 不在精选清单,必须被排除
|
||||
})
|
||||
|
||||
count, created, err := svc.EnsureDefaultEmojiStatusSet(ctx)
|
||||
if err != nil {
|
||||
t.Fatalf("ensure: %v", err)
|
||||
}
|
||||
if !created || count != 3 {
|
||||
t.Fatalf("ensure = (count=%d, created=%v), want (3, true)", count, created)
|
||||
}
|
||||
set, found, err := media.GetStickerSetBySystemKey(ctx, domain.StickerSetSystemKeyEmojiDefaultStatuses)
|
||||
if err != nil || !found {
|
||||
t.Fatalf("synthesized set not found: found=%v err=%v", found, err)
|
||||
}
|
||||
if set.Kind != domain.StickerSetKindSystem || !set.Emojis || set.Count != 3 || set.Hash == 0 {
|
||||
t.Fatalf("set meta = %+v, want system kind, emojis, count 3, non-zero hash", set)
|
||||
}
|
||||
got := map[int64]bool{}
|
||||
for _, id := range set.DocumentIDs {
|
||||
got[id] = true
|
||||
}
|
||||
if !got[101] || !got[102] || !got[103] || got[999] {
|
||||
t.Fatalf("document ids = %v, want 101/102/103 without 999", set.DocumentIDs)
|
||||
}
|
||||
// 精选顺序:☕(=103) 在 ❤(=101) 之前、❤ 在 👍(=102) 之前(按清单序而非 pack 序)。
|
||||
index := map[int64]int{}
|
||||
for i, id := range set.DocumentIDs {
|
||||
index[id] = i
|
||||
}
|
||||
if !(index[103] < index[101] && index[101] < index[102]) {
|
||||
t.Fatalf("document order = %v, want curated order ☕<❤<👍", set.DocumentIDs)
|
||||
}
|
||||
|
||||
// 幂等:第二次调用不得重建。
|
||||
count2, created2, err := svc.EnsureDefaultEmojiStatusSet(ctx)
|
||||
if err != nil {
|
||||
t.Fatalf("ensure again: %v", err)
|
||||
}
|
||||
if created2 || count2 != 3 {
|
||||
t.Fatalf("ensure again = (count=%d, created=%v), want (3, false)", count2, created2)
|
||||
}
|
||||
|
||||
// ResolveStickerSet(inputStickerSetEmojiDefaultStatuses 的服务路径)能解析。
|
||||
resolved, docs, found, err := svc.ResolveStickerSet(ctx, domain.StickerSetRef{
|
||||
Kind: domain.StickerSetRefBySystem,
|
||||
SystemKey: domain.StickerSetSystemKeyEmojiDefaultStatuses,
|
||||
})
|
||||
if err != nil || !found {
|
||||
t.Fatalf("resolve: found=%v err=%v", found, err)
|
||||
}
|
||||
if resolved.ID != set.ID || len(docs) != 3 {
|
||||
t.Fatalf("resolve = set %d with %d docs, want %d with 3", resolved.ID, len(docs), set.ID)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEnsureDefaultEmojiStatusSetWithoutSeed(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
svc, media := defaultStatusesTestService(t)
|
||||
count, created, err := svc.EnsureDefaultEmojiStatusSet(ctx)
|
||||
if err != nil {
|
||||
t.Fatalf("ensure without seed: %v", err)
|
||||
}
|
||||
if created || count != 0 {
|
||||
t.Fatalf("ensure without seed = (count=%d, created=%v), want (0, false)", count, created)
|
||||
}
|
||||
if _, found, _ := media.GetStickerSetBySystemKey(ctx, domain.StickerSetSystemKeyEmojiDefaultStatuses); found {
|
||||
t.Fatal("set must not be created without animated_emoji seed")
|
||||
}
|
||||
}
|
||||
167
internal/app/files/effects_seed.go
Normal file
167
internal/app/files/effects_seed.go
Normal file
|
|
@ -0,0 +1,167 @@
|
|||
package files
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/binary"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"hash"
|
||||
"hash/fnv"
|
||||
"os"
|
||||
"path/filepath"
|
||||
|
||||
"telesrv/internal/domain"
|
||||
)
|
||||
|
||||
// telegram_effects_export/effects.json 的解析结构。effects 引用文档 id,文档全量元数据
|
||||
// 在 documents[] 里(与 messages.availableEffects 同构),blob 在 documents/<docid>.<ext>。
|
||||
type seedEffectJSON struct {
|
||||
ID int64 `json:"id"`
|
||||
Emoticon string `json:"emoticon"`
|
||||
StaticIconID int64 `json:"static_icon_id"`
|
||||
EffectStickerID int64 `json:"effect_sticker_id"`
|
||||
EffectAnimationID int64 `json:"effect_animation_id"`
|
||||
PremiumRequired bool `json:"premium_required"`
|
||||
}
|
||||
|
||||
type seedEffectsFileJSON struct {
|
||||
Result struct {
|
||||
Effects []seedEffectJSON `json:"effects"`
|
||||
Documents []seedDocumentJSON `json:"documents"`
|
||||
} `json:"result"`
|
||||
}
|
||||
|
||||
// seedEffects 从 telegram_effects_export 导入消息特效。特效元数据每次启动都从 JSON
|
||||
// 重建进内存;document/blob 只有 catalog hash 变化或持久化资源不完整时才重导。
|
||||
func (s *Service) seedEffects(ctx context.Context, root string, stats *SeedStats) error {
|
||||
dir := filepath.Join(root, "telegram_effects_export")
|
||||
raw, err := os.ReadFile(filepath.Join(dir, "effects.json"))
|
||||
if err != nil {
|
||||
return nil // 无 effects 资源 → 跳过
|
||||
}
|
||||
var parsed seedEffectsFileJSON
|
||||
if err := json.Unmarshal(raw, &parsed); err != nil {
|
||||
return fmt.Errorf("parse effects.json: %w", err)
|
||||
}
|
||||
docsDir := filepath.Join(dir, "documents")
|
||||
index, err := scanSeedDir(docsDir)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
effects, requiredDocs := seedEffectsCatalog(parsed)
|
||||
s.effects = effects
|
||||
s.effectsHash = effectsCatalogHash(effects)
|
||||
stats.Effects = len(effects)
|
||||
|
||||
stateHash, err := s.seedEffectsStateHash(raw, docsDir)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
ready, err := s.seedDocumentJSONsReady(ctx, requiredDocs, index)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
matched, err := s.seedStateMatches(ctx, seedEffectsStateKey, stateHash)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if matched && ready {
|
||||
return nil
|
||||
}
|
||||
|
||||
// 多个 effect 常共享同一文档(static icon 尤甚):每个唯一源文档只导一次。
|
||||
for _, dj := range requiredDocs {
|
||||
if _, err := s.importDocument(ctx, dj, docsDir, index, stats); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return s.putSeedState(ctx, seedEffectsStateKey, stateHash)
|
||||
}
|
||||
|
||||
func seedEffectsCatalog(parsed seedEffectsFileJSON) ([]domain.AvailableEffect, []seedDocumentJSON) {
|
||||
docByID := make(map[int64]seedDocumentJSON, len(parsed.Result.Documents))
|
||||
for _, d := range parsed.Result.Documents {
|
||||
docByID[d.ID] = d
|
||||
}
|
||||
required := make(map[int64]struct{}, len(docByID))
|
||||
storageID := func(sourceID int64) int64 {
|
||||
if sourceID == 0 {
|
||||
return 0
|
||||
}
|
||||
if _, ok := docByID[sourceID]; !ok {
|
||||
return 0
|
||||
}
|
||||
required[sourceID] = struct{}{}
|
||||
return seedDocumentStorageID(sourceID)
|
||||
}
|
||||
effects := make([]domain.AvailableEffect, 0, len(parsed.Result.Effects))
|
||||
for i, ej := range parsed.Result.Effects {
|
||||
if ej.ID == 0 || ej.EffectStickerID == 0 {
|
||||
continue
|
||||
}
|
||||
staticID := storageID(ej.StaticIconID)
|
||||
stickerID := storageID(ej.EffectStickerID)
|
||||
if stickerID == 0 {
|
||||
continue
|
||||
}
|
||||
animID := storageID(ej.EffectAnimationID)
|
||||
effects = append(effects, domain.AvailableEffect{
|
||||
ID: ej.ID,
|
||||
Emoticon: ej.Emoticon,
|
||||
StaticIconID: staticID,
|
||||
EffectStickerID: stickerID,
|
||||
EffectAnimationID: animID,
|
||||
PremiumRequired: ej.PremiumRequired,
|
||||
Order: i,
|
||||
})
|
||||
}
|
||||
docs := make([]seedDocumentJSON, 0, len(required))
|
||||
for _, d := range parsed.Result.Documents {
|
||||
if _, ok := required[d.ID]; ok {
|
||||
docs = append(docs, d)
|
||||
}
|
||||
}
|
||||
return effects, docs
|
||||
}
|
||||
|
||||
func (s *Service) seedEffectsStateHash(raw []byte, docsDir string) (string, error) {
|
||||
return seedStateHash(func(h hash.Hash) error {
|
||||
writeSeedStateHeader(h, seedEffectsStateVersion, s.dc)
|
||||
_, _ = h.Write(raw)
|
||||
_, _ = h.Write([]byte{'\n'})
|
||||
return writeSeedDirFingerprint(h, docsDir)
|
||||
})
|
||||
}
|
||||
|
||||
// AvailableEffects 返回 seed 进内存的消息特效目录与其内容 hash(全局静态;hash 在 seed 时
|
||||
// 算一次,handler 直接比对返回 NotModified,无需每次 RPC 重算)。
|
||||
func (s *Service) AvailableEffects(ctx context.Context) ([]domain.AvailableEffect, int, error) {
|
||||
return s.effects, s.effectsHash, nil
|
||||
}
|
||||
|
||||
// effectsCatalogHash 由 effect 字段算稳定正整数 hash(FNV-1a)。内容变则 hash 变,
|
||||
// 客户端发旧 hash 即不命中→重取。
|
||||
func effectsCatalogHash(effects []domain.AvailableEffect) int {
|
||||
if len(effects) == 0 {
|
||||
return 0
|
||||
}
|
||||
h := fnv.New64a()
|
||||
var buf [8]byte
|
||||
put := func(v int64) {
|
||||
binary.LittleEndian.PutUint64(buf[:], uint64(v))
|
||||
_, _ = h.Write(buf[:])
|
||||
}
|
||||
for _, e := range effects {
|
||||
put(e.ID)
|
||||
_, _ = h.Write([]byte(e.Emoticon))
|
||||
put(e.StaticIconID)
|
||||
put(e.EffectStickerID)
|
||||
put(e.EffectAnimationID)
|
||||
if e.PremiumRequired {
|
||||
put(1)
|
||||
} else {
|
||||
put(0)
|
||||
}
|
||||
}
|
||||
return int(h.Sum64() & 0x7fffffff)
|
||||
}
|
||||
50
internal/app/files/encrypted_file.go
Normal file
50
internal/app/files/encrypted_file.go
Normal file
|
|
@ -0,0 +1,50 @@
|
|||
package files
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
|
||||
"go.uber.org/zap"
|
||||
|
||||
"telesrv/internal/domain"
|
||||
)
|
||||
|
||||
// CreateEncryptedFileFromUpload 把已上传分片组装成密聊文件 blob 并铸造 EncryptedFile 快照。
|
||||
// 盲中继:内容是客户端加密的 bytes,不解析、不分类、不缩略图;blob 落 location_key
|
||||
// "enc:<id>"(复用 BlobBackend,下载经 inputEncryptedFileLocation → 同 key)。
|
||||
// access_hash 不强校验(沿用现有媒体 dev 姿态,依赖不可枚举 id)。元数据持久化由调用方
|
||||
// (rpc 层经 SecretChats.PutEncryptedFile)负责。
|
||||
func (s *Service) CreateEncryptedFileFromUpload(ctx context.Context, file domain.UploadedFileRef, keyFingerprint int) (domain.EncryptedFileRef, error) {
|
||||
body, err := s.assembleUploadBlob(ctx, file.OwnerUserID, file.FileID, file.Parts)
|
||||
if err != nil {
|
||||
return domain.EncryptedFileRef{}, err
|
||||
}
|
||||
if body.Size == 0 {
|
||||
return domain.EncryptedFileRef{}, domain.ErrDocumentInvalid
|
||||
}
|
||||
id := randomID()
|
||||
if err := s.media.PutFileBlob(ctx, domain.FileBlob{
|
||||
LocationKey: fmt.Sprintf("enc:%d", id),
|
||||
Backend: domain.MediaBackend(s.blobs.Name()),
|
||||
ObjectKey: body.ObjectKey,
|
||||
Size: body.Size,
|
||||
SHA256: body.SHA256,
|
||||
MimeType: "application/octet-stream",
|
||||
}); err != nil {
|
||||
return domain.EncryptedFileRef{}, err
|
||||
}
|
||||
if err := s.cleanupUploadParts(ctx, file.OwnerUserID, file.FileID); err != nil {
|
||||
s.log.Warn("cleanup encrypted file upload parts failed",
|
||||
zap.Int64("owner_user_id", file.OwnerUserID),
|
||||
zap.Int64("file_id", file.FileID),
|
||||
zap.Int64("encrypted_file_id", id),
|
||||
zap.Error(err))
|
||||
}
|
||||
return domain.EncryptedFileRef{
|
||||
ID: id,
|
||||
AccessHash: randomID(),
|
||||
Size: body.Size,
|
||||
DCID: s.dc,
|
||||
KeyFingerprint: keyFingerprint,
|
||||
}, nil
|
||||
}
|
||||
223
internal/app/files/external_media.go
Normal file
223
internal/app/files/external_media.go
Normal file
|
|
@ -0,0 +1,223 @@
|
|||
package files
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"path"
|
||||
"strings"
|
||||
"sync"
|
||||
"syscall"
|
||||
"time"
|
||||
|
||||
"telesrv/internal/domain"
|
||||
)
|
||||
|
||||
// 外链媒体:inputMediaPhotoExternal / inputMediaDocumentExternal——客户端给一个 URL,
|
||||
// 服务端抓取并铸造 Photo/Document。抓取任意用户可控 URL,安全是核心:
|
||||
// - SSRF 防护:自定义 Dialer.Control 在连接前检查**解析出的目标 IP**,挡掉 loopback/
|
||||
// 私网/link-local/CGNAT/multicast/unspecified。因为每次实际 dial 都查,所以同时防住
|
||||
// DNS rebinding(公网域名解析到内网 IP)与重定向(每一跳都重新 dial→重新检查)。
|
||||
// - 仅 http/https;重定向上限;响应大小上限(LimitReader);请求超时;全局抓取限速
|
||||
// (防一条消息触发大量服务端外网抓取的放大攻击)。
|
||||
|
||||
var (
|
||||
// ErrExternalMediaDisabled 表示未启用外链媒体抓取(rpc 层映射为 MEDIA_INVALID)。
|
||||
ErrExternalMediaDisabled = errors.New("external media disabled")
|
||||
// ErrExternalMediaInvalid 表示 URL 不合法/被 SSRF 防护拦截/上游失败/超限。
|
||||
ErrExternalMediaInvalid = errors.New("external media invalid")
|
||||
)
|
||||
|
||||
const (
|
||||
externalMediaTimeout = 15 * time.Second
|
||||
externalMediaMaxRedirects = 5
|
||||
// DefaultExternalMediaMaxBytes 是抓取响应体上限。
|
||||
DefaultExternalMediaMaxBytes = int64(10 << 20)
|
||||
// DefaultExternalMediaRatePerMin 是全局每分钟抓取上限(防放大攻击)。
|
||||
DefaultExternalMediaRatePerMin = 60
|
||||
externalMediaRateWindow = time.Minute
|
||||
)
|
||||
|
||||
type externalMediaFetcher struct {
|
||||
client *http.Client
|
||||
maxBytes int64
|
||||
rateLimit int
|
||||
|
||||
mu sync.Mutex
|
||||
fetchTimes []time.Time
|
||||
}
|
||||
|
||||
// WithExternalMedia 启用外链媒体抓取(inputMediaPhoto/DocumentExternal)。
|
||||
// maxBytes<=0 用默认;ratePerMin<=0 用默认。SSRF 防护恒开。
|
||||
func WithExternalMedia(maxBytes int64, ratePerMin int) Option {
|
||||
return func(s *Service) {
|
||||
if maxBytes <= 0 {
|
||||
maxBytes = DefaultExternalMediaMaxBytes
|
||||
}
|
||||
if ratePerMin <= 0 {
|
||||
ratePerMin = DefaultExternalMediaRatePerMin
|
||||
}
|
||||
s.externalMedia = newExternalMediaFetcher(maxBytes, ratePerMin, false)
|
||||
}
|
||||
}
|
||||
|
||||
// newExternalMediaFetcher 构造抓取器。allowPrivate 仅供测试(指向 httptest loopback);
|
||||
// 生产恒 false。
|
||||
func newExternalMediaFetcher(maxBytes int64, ratePerMin int, allowPrivate bool) *externalMediaFetcher {
|
||||
dialer := &net.Dialer{Timeout: externalMediaTimeout}
|
||||
dialer.Control = func(network, address string, _ syscall.RawConn) error {
|
||||
host, _, err := net.SplitHostPort(address)
|
||||
if err != nil {
|
||||
return ErrExternalMediaInvalid
|
||||
}
|
||||
ip := net.ParseIP(host)
|
||||
if ip == nil {
|
||||
return ErrExternalMediaInvalid
|
||||
}
|
||||
if !allowPrivate && isBlockedExternalIP(ip) {
|
||||
return fmt.Errorf("%w: blocked address %s (SSRF guard)", ErrExternalMediaInvalid, host)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
client := &http.Client{
|
||||
Timeout: externalMediaTimeout,
|
||||
Transport: &http.Transport{DialContext: dialer.DialContext, DisableKeepAlives: true},
|
||||
CheckRedirect: func(req *http.Request, via []*http.Request) error {
|
||||
if len(via) >= externalMediaMaxRedirects {
|
||||
return fmt.Errorf("%w: too many redirects", ErrExternalMediaInvalid)
|
||||
}
|
||||
if req.URL.Scheme != "http" && req.URL.Scheme != "https" {
|
||||
return fmt.Errorf("%w: blocked redirect scheme %q", ErrExternalMediaInvalid, req.URL.Scheme)
|
||||
}
|
||||
return nil
|
||||
},
|
||||
}
|
||||
return &externalMediaFetcher{client: client, maxBytes: maxBytes, rateLimit: ratePerMin}
|
||||
}
|
||||
|
||||
// isBlockedExternalIP 报告是否为不可对外抓取的内网/特殊地址(SSRF 防护)。
|
||||
func isBlockedExternalIP(ip net.IP) bool {
|
||||
if ip.IsLoopback() || ip.IsPrivate() || ip.IsUnspecified() ||
|
||||
ip.IsLinkLocalUnicast() || ip.IsLinkLocalMulticast() ||
|
||||
ip.IsMulticast() || ip.IsInterfaceLocalMulticast() {
|
||||
return true
|
||||
}
|
||||
// CGNAT 100.64.0.0/10(运营商级 NAT,常用于内部基础设施)。
|
||||
if ip4 := ip.To4(); ip4 != nil && ip4[0] == 100 && ip4[1] >= 64 && ip4[1] <= 127 {
|
||||
return true
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func (f *externalMediaFetcher) allowFetch() bool {
|
||||
f.mu.Lock()
|
||||
defer f.mu.Unlock()
|
||||
now := time.Now()
|
||||
kept := f.fetchTimes[:0]
|
||||
for _, at := range f.fetchTimes {
|
||||
if now.Sub(at) <= externalMediaRateWindow {
|
||||
kept = append(kept, at)
|
||||
}
|
||||
}
|
||||
f.fetchTimes = kept
|
||||
if len(f.fetchTimes) >= f.rateLimit {
|
||||
return false
|
||||
}
|
||||
f.fetchTimes = append(f.fetchTimes, now)
|
||||
return true
|
||||
}
|
||||
|
||||
// fetch 抓取 URL,返回 (字节, content-type)。SSRF 检查在 dial 阶段发生。
|
||||
func (f *externalMediaFetcher) fetch(ctx context.Context, rawURL string) ([]byte, string, error) {
|
||||
u, err := url.Parse(strings.TrimSpace(rawURL))
|
||||
if err != nil || (u.Scheme != "http" && u.Scheme != "https") || u.Host == "" {
|
||||
return nil, "", ErrExternalMediaInvalid
|
||||
}
|
||||
if !f.allowFetch() {
|
||||
return nil, "", fmt.Errorf("%w: rate limited", ErrExternalMediaInvalid)
|
||||
}
|
||||
ctx, cancel := context.WithTimeout(ctx, externalMediaTimeout)
|
||||
defer cancel()
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, u.String(), nil)
|
||||
if err != nil {
|
||||
return nil, "", ErrExternalMediaInvalid
|
||||
}
|
||||
req.Header.Set("User-Agent", "telesrv-media-fetch")
|
||||
resp, err := f.client.Do(req)
|
||||
if err != nil {
|
||||
// 含 SSRF 拦截、超时、传输错误。
|
||||
return nil, "", fmt.Errorf("%w: %v", ErrExternalMediaInvalid, err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return nil, "", fmt.Errorf("%w: upstream status %d", ErrExternalMediaInvalid, resp.StatusCode)
|
||||
}
|
||||
data, err := io.ReadAll(io.LimitReader(resp.Body, f.maxBytes+1))
|
||||
if err != nil {
|
||||
return nil, "", fmt.Errorf("%w: read body: %v", ErrExternalMediaInvalid, err)
|
||||
}
|
||||
if len(data) == 0 || int64(len(data)) > f.maxBytes {
|
||||
return nil, "", fmt.Errorf("%w: body size %d", ErrExternalMediaInvalid, len(data))
|
||||
}
|
||||
contentType := resp.Header.Get("Content-Type")
|
||||
if i := strings.IndexByte(contentType, ';'); i >= 0 {
|
||||
contentType = contentType[:i]
|
||||
}
|
||||
return data, strings.TrimSpace(contentType), nil
|
||||
}
|
||||
|
||||
// CreatePhotoFromURL 抓取 URL 并铸造 Photo(CreatePhotoFromBytes 会解码校验是否为图片)。
|
||||
func (s *Service) CreatePhotoFromURL(ctx context.Context, rawURL string) (domain.Photo, error) {
|
||||
if s == nil || s.externalMedia == nil {
|
||||
return domain.Photo{}, ErrExternalMediaDisabled
|
||||
}
|
||||
data, _, err := s.externalMedia.fetch(ctx, rawURL)
|
||||
if err != nil {
|
||||
return domain.Photo{}, err
|
||||
}
|
||||
photo, err := s.CreatePhotoFromBytes(ctx, data)
|
||||
if err != nil {
|
||||
// 非图片字节 → ErrPhotoInvalid,对外统一为 external invalid。
|
||||
return domain.Photo{}, fmt.Errorf("%w: %v", ErrExternalMediaInvalid, err)
|
||||
}
|
||||
return photo, nil
|
||||
}
|
||||
|
||||
// CreateDocumentFromURL 抓取 URL 并铸造 Document:mime 取 Content-Type,文件名取 URL basename。
|
||||
func (s *Service) CreateDocumentFromURL(ctx context.Context, rawURL string) (domain.Document, error) {
|
||||
if s == nil || s.externalMedia == nil {
|
||||
return domain.Document{}, ErrExternalMediaDisabled
|
||||
}
|
||||
data, contentType, err := s.externalMedia.fetch(ctx, rawURL)
|
||||
if err != nil {
|
||||
return domain.Document{}, err
|
||||
}
|
||||
mime := contentType
|
||||
if mime == "" {
|
||||
mime = "application/octet-stream"
|
||||
}
|
||||
spec := domain.DocumentSpec{
|
||||
MimeType: mime,
|
||||
Attributes: []domain.DocumentAttribute{{Kind: domain.DocAttrFilename, FileName: externalMediaFilename(rawURL)}},
|
||||
}
|
||||
doc, err := s.CreateDocumentFromBytes(ctx, data, spec)
|
||||
if err != nil {
|
||||
return domain.Document{}, fmt.Errorf("%w: %v", ErrExternalMediaInvalid, err)
|
||||
}
|
||||
return doc, nil
|
||||
}
|
||||
|
||||
// externalMediaFilename 从 URL path 取 basename;缺失时回退通用名。
|
||||
func externalMediaFilename(rawURL string) string {
|
||||
u, err := url.Parse(rawURL)
|
||||
if err == nil {
|
||||
if base := path.Base(u.Path); base != "" && base != "." && base != "/" {
|
||||
return base
|
||||
}
|
||||
}
|
||||
return "file"
|
||||
}
|
||||
106
internal/app/files/external_media_test.go
Normal file
106
internal/app/files/external_media_test.go
Normal file
|
|
@ -0,0 +1,106 @@
|
|||
package files
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestIsBlockedExternalIP(t *testing.T) {
|
||||
cases := []struct {
|
||||
ip string
|
||||
blocked bool
|
||||
}{
|
||||
{"127.0.0.1", true}, // loopback
|
||||
{"::1", true}, // loopback v6
|
||||
{"10.0.0.5", true}, // private
|
||||
{"172.16.3.4", true}, // private
|
||||
{"192.168.1.1", true}, // private
|
||||
{"169.254.1.1", true}, // link-local
|
||||
{"fe80::1", true}, // link-local v6
|
||||
{"0.0.0.0", true}, // unspecified
|
||||
{"100.64.0.1", true}, // CGNAT
|
||||
{"100.127.255.1", true}, // CGNAT 上界
|
||||
{"224.0.0.1", true}, // multicast
|
||||
{"8.8.8.8", false}, // 公网
|
||||
{"1.1.1.1", false}, // 公网
|
||||
{"100.63.255.1", false}, // CGNAT 下界外(公网)
|
||||
{"100.128.0.1", false}, // CGNAT 上界外(公网)
|
||||
{"2606:4700:4700::1111", false}, // 公网 v6
|
||||
}
|
||||
for _, c := range cases {
|
||||
ip := net.ParseIP(c.ip)
|
||||
if ip == nil {
|
||||
t.Fatalf("parse %s failed", c.ip)
|
||||
}
|
||||
if got := isBlockedExternalIP(ip); got != c.blocked {
|
||||
t.Errorf("isBlockedExternalIP(%s) = %v, want %v", c.ip, got, c.blocked)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestExternalMediaFetcherSSRFGuard 验证 SSRF 防护:httptest 在 loopback 上,
|
||||
// allowPrivate=false 必须拦截(不连接内网),allowPrivate=true 放行抓取到字节。
|
||||
func TestExternalMediaFetcherSSRFGuard(t *testing.T) {
|
||||
body := []byte("hello-external-bytes")
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/octet-stream")
|
||||
_, _ = w.Write(body)
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
// 生产配置(allowPrivate=false):loopback 目标被 SSRF 防护拦截。
|
||||
guarded := newExternalMediaFetcher(DefaultExternalMediaMaxBytes, DefaultExternalMediaRatePerMin, false)
|
||||
if _, _, err := guarded.fetch(context.Background(), srv.URL); !errors.Is(err, ErrExternalMediaInvalid) {
|
||||
t.Fatalf("SSRF guard fetch err = %v, want ErrExternalMediaInvalid (loopback 应被拦)", err)
|
||||
}
|
||||
|
||||
// 测试放行(allowPrivate=true):抓取成功。
|
||||
open := newExternalMediaFetcher(DefaultExternalMediaMaxBytes, DefaultExternalMediaRatePerMin, true)
|
||||
data, ct, err := open.fetch(context.Background(), srv.URL)
|
||||
if err != nil {
|
||||
t.Fatalf("open fetch err = %v", err)
|
||||
}
|
||||
if string(data) != string(body) {
|
||||
t.Fatalf("fetched %q, want %q", data, body)
|
||||
}
|
||||
if ct != "application/octet-stream" {
|
||||
t.Fatalf("content-type = %q, want application/octet-stream", ct)
|
||||
}
|
||||
}
|
||||
|
||||
// TestExternalMediaFetcherRejectsBadURL 非 http(s)/空 host 直接拒。
|
||||
func TestExternalMediaFetcherRejectsBadURL(t *testing.T) {
|
||||
f := newExternalMediaFetcher(DefaultExternalMediaMaxBytes, DefaultExternalMediaRatePerMin, true)
|
||||
for _, bad := range []string{"", "ftp://x/y", "file:///etc/passwd", "javascript:alert(1)", "http://", "not a url"} {
|
||||
if _, _, err := f.fetch(context.Background(), bad); !errors.Is(err, ErrExternalMediaInvalid) {
|
||||
t.Errorf("fetch(%q) err = %v, want ErrExternalMediaInvalid", bad, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestExternalMediaFetcherSizeLimit 超大小上限拒。
|
||||
func TestExternalMediaFetcherSizeLimit(t *testing.T) {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
_, _ = w.Write(make([]byte, 2048))
|
||||
}))
|
||||
defer srv.Close()
|
||||
f := newExternalMediaFetcher(1024, DefaultExternalMediaRatePerMin, true)
|
||||
if _, _, err := f.fetch(context.Background(), srv.URL); !errors.Is(err, ErrExternalMediaInvalid) {
|
||||
t.Fatalf("oversize fetch err = %v, want ErrExternalMediaInvalid", err)
|
||||
}
|
||||
}
|
||||
|
||||
// TestExternalMediaDisabled 未启用时 Create*FromURL 返回 ErrExternalMediaDisabled。
|
||||
func TestExternalMediaDisabled(t *testing.T) {
|
||||
s := &Service{}
|
||||
if _, err := s.CreatePhotoFromURL(context.Background(), "http://x/y.png"); !errors.Is(err, ErrExternalMediaDisabled) {
|
||||
t.Fatalf("disabled photo err = %v, want ErrExternalMediaDisabled", err)
|
||||
}
|
||||
if _, err := s.CreateDocumentFromURL(context.Background(), "http://x/y.bin"); !errors.Is(err, ErrExternalMediaDisabled) {
|
||||
t.Fatalf("disabled doc err = %v, want ErrExternalMediaDisabled", err)
|
||||
}
|
||||
}
|
||||
183
internal/app/files/maptile.go
Normal file
183
internal/app/files/maptile.go
Normal file
|
|
@ -0,0 +1,183 @@
|
|||
package files
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"image"
|
||||
"image/color"
|
||||
"image/png"
|
||||
"math"
|
||||
|
||||
"go.uber.org/zap"
|
||||
)
|
||||
|
||||
// 本文件实现 geo 消息地图缩略图(upload.getWebFile / inputWebFileGeoPointLocation)的
|
||||
// 本地占位渲染:服务端按坐标确定性合成一张「街区网格 + 定位针」风格的静态图,
|
||||
// 同一 (lat,long,zoom,w,h,scale) 输入字节级可重现,保证客户端分片续传一致。
|
||||
// 配置了 Mapbox 代理(maptile_proxy.go)时优先返回真实地图,本渲染降级为故障回退。
|
||||
|
||||
const (
|
||||
mapTileMinEdge = 16
|
||||
mapTileMaxEdge = 1024
|
||||
mapTileMinZoom = 13
|
||||
mapTileMaxZoom = 20
|
||||
mapTileMaxScale = 3
|
||||
)
|
||||
|
||||
// GeoMapTile 返回一张 w×h(逻辑像素,输出按 scale 放大)的静态地图与 mime。
|
||||
// 入参越界时按协议约束 clamp(w/h 16-1024、zoom 13-20、scale 1-3),不报错。
|
||||
// 配置了 Mapbox 代理时优先真实地图(落盘缓存),抓取失败回退确定性占位渲染。
|
||||
func (s *Service) GeoMapTile(lat, long float64, w, h, zoom, scale int) ([]byte, string) {
|
||||
w = clampInt(w, mapTileMinEdge, mapTileMaxEdge)
|
||||
h = clampInt(h, mapTileMinEdge, mapTileMaxEdge)
|
||||
zoom = clampInt(zoom, mapTileMinZoom, mapTileMaxZoom)
|
||||
scale = clampInt(scale, 1, mapTileMaxScale)
|
||||
// nil receiver 合法:占位渲染是纯函数(既有调用方/测试依赖这一点),代理仅在配置后启用。
|
||||
if s != nil && s.mapTiles != nil {
|
||||
if data, mime, err := s.mapTiles.tile(lat, long, w, h, zoom, scale); err == nil {
|
||||
return data, mime
|
||||
} else if s.log != nil {
|
||||
s.log.Warn("map tile proxy failed, fallback to placeholder",
|
||||
zap.Error(err), zap.Float64("lat", lat), zap.Float64("long", long), zap.Int("zoom", zoom))
|
||||
}
|
||||
}
|
||||
pw, ph := w*scale, h*scale
|
||||
|
||||
img := image.NewRGBA(image.Rect(0, 0, pw, ph))
|
||||
background := color.RGBA{R: 0xEB, G: 0xE7, B: 0xDE, A: 0xFF}
|
||||
park := color.RGBA{R: 0xCF, G: 0xE4, B: 0xC2, A: 0xFF}
|
||||
road := color.RGBA{R: 0xFF, G: 0xFF, B: 0xFF, A: 0xFF}
|
||||
roadEdge := color.RGBA{R: 0xD9, G: 0xD3, B: 0xC7, A: 0xFF}
|
||||
fillRect(img, 0, 0, pw, ph, background)
|
||||
|
||||
rng := newMapTileRNG(lat, long, zoom)
|
||||
|
||||
// 街区网格:间距与抖动由坐标种子决定,平移随经纬度连续变化,避免所有地点一张脸。
|
||||
spacing := (48 + int(rng.next()%32)) * scale
|
||||
offX := int(math.Abs(long*1e4)) % spacing
|
||||
offY := int(math.Abs(lat*1e4)) % spacing
|
||||
for x := -offX; x < pw; x += spacing {
|
||||
major := ((x+offX)/spacing)%3 == int(rng.next()%3)
|
||||
drawVerticalRoad(img, x+int(rng.next()%uint64(spacing/3)), pw, ph, scale, major, road, roadEdge)
|
||||
}
|
||||
for y := -offY; y < ph; y += spacing {
|
||||
major := ((y+offY)/spacing)%3 == int(rng.next()%3)
|
||||
drawHorizontalRoad(img, y+int(rng.next()%uint64(spacing/3)), pw, ph, scale, major, road, roadEdge)
|
||||
}
|
||||
|
||||
// 两块「绿地」:取网格内随机街块,铺底色之上、道路之下的视觉层级太复杂,
|
||||
// 这里直接半覆盖即可(占位图不追求制图精度)。
|
||||
for i := 0; i < 2; i++ {
|
||||
bx := int(rng.next() % uint64(pw))
|
||||
by := int(rng.next() % uint64(ph))
|
||||
bw := (spacing * 3) / 4
|
||||
fillRect(img, bx, by, minInt(bx+bw, pw), minInt(by+bw, ph), park)
|
||||
}
|
||||
|
||||
drawCenterPin(img, pw, ph, scale)
|
||||
|
||||
var buf bytes.Buffer
|
||||
_ = png.Encode(&buf, img)
|
||||
return buf.Bytes(), "image/png"
|
||||
}
|
||||
|
||||
// mapTileRNG 是确定性 xorshift64,种子来自量化坐标与 zoom。
|
||||
type mapTileRNG struct{ state uint64 }
|
||||
|
||||
func newMapTileRNG(lat, long float64, zoom int) *mapTileRNG {
|
||||
seed := uint64(int64(lat*1e5))*1000003 ^ uint64(int64(long*1e5))*998244353 ^ uint64(zoom)*0x9E3779B97F4A7C15
|
||||
if seed == 0 {
|
||||
seed = 0x9E3779B97F4A7C15
|
||||
}
|
||||
return &mapTileRNG{state: seed}
|
||||
}
|
||||
|
||||
func (r *mapTileRNG) next() uint64 {
|
||||
r.state ^= r.state << 13
|
||||
r.state ^= r.state >> 7
|
||||
r.state ^= r.state << 17
|
||||
return r.state
|
||||
}
|
||||
|
||||
func fillRect(img *image.RGBA, x0, y0, x1, y1 int, c color.RGBA) {
|
||||
bounds := img.Bounds()
|
||||
x0, y0 = maxInt(x0, bounds.Min.X), maxInt(y0, bounds.Min.Y)
|
||||
x1, y1 = minInt(x1, bounds.Max.X), minInt(y1, bounds.Max.Y)
|
||||
for y := y0; y < y1; y++ {
|
||||
for x := x0; x < x1; x++ {
|
||||
img.SetRGBA(x, y, c)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func drawVerticalRoad(img *image.RGBA, x, pw, ph, scale int, major bool, road, edge color.RGBA) {
|
||||
width := 2 * scale
|
||||
if major {
|
||||
width = 4 * scale
|
||||
}
|
||||
fillRect(img, x-width/2-1, 0, x+width/2+1, ph, edge)
|
||||
fillRect(img, x-width/2, 0, x+width/2, ph, road)
|
||||
}
|
||||
|
||||
func drawHorizontalRoad(img *image.RGBA, y, pw, ph, scale int, major bool, road, edge color.RGBA) {
|
||||
width := 2 * scale
|
||||
if major {
|
||||
width = 4 * scale
|
||||
}
|
||||
fillRect(img, 0, y-width/2-1, pw, y+width/2+1, edge)
|
||||
fillRect(img, 0, y-width/2, pw, y+width/2, road)
|
||||
}
|
||||
|
||||
// drawCenterPin 在图中心画红色定位针(圆头 + 下尖三角 + 白色内点),中心即坐标点。
|
||||
func drawCenterPin(img *image.RGBA, pw, ph, scale int) {
|
||||
pin := color.RGBA{R: 0xE5, G: 0x39, B: 0x35, A: 0xFF}
|
||||
pinDark := color.RGBA{R: 0xB7, G: 0x1C, B: 0x1C, A: 0xFF}
|
||||
white := color.RGBA{R: 0xFF, G: 0xFF, B: 0xFF, A: 0xFF}
|
||||
cx, cy := pw/2, ph/2
|
||||
headR := 9 * scale
|
||||
headCY := cy - 14*scale
|
||||
// 尖角三角形:从圆头两侧收敛到坐标点。
|
||||
for y := headCY; y <= cy; y++ {
|
||||
t := float64(y-headCY) / float64(cy-headCY)
|
||||
half := int(float64(headR) * (1 - t) * 0.82)
|
||||
for x := cx - half; x <= cx+half; x++ {
|
||||
img.SetRGBA(x, y, pin)
|
||||
}
|
||||
}
|
||||
// 圆头(带一圈深色描边)。
|
||||
for dy := -headR - scale; dy <= headR+scale; dy++ {
|
||||
for dx := -headR - scale; dx <= headR+scale; dx++ {
|
||||
d2 := dx*dx + dy*dy
|
||||
switch {
|
||||
case d2 <= headR*headR:
|
||||
img.SetRGBA(cx+dx, headCY+dy, pin)
|
||||
case d2 <= (headR+scale)*(headR+scale):
|
||||
img.SetRGBA(cx+dx, headCY+dy, pinDark)
|
||||
}
|
||||
}
|
||||
}
|
||||
innerR := 3 * scale
|
||||
for dy := -innerR; dy <= innerR; dy++ {
|
||||
for dx := -innerR; dx <= innerR; dx++ {
|
||||
if dx*dx+dy*dy <= innerR*innerR {
|
||||
img.SetRGBA(cx+dx, headCY+dy, white)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func clampInt(v, lo, hi int) int {
|
||||
if v < lo {
|
||||
return lo
|
||||
}
|
||||
if v > hi {
|
||||
return hi
|
||||
}
|
||||
return v
|
||||
}
|
||||
|
||||
func minInt(a, b int) int {
|
||||
if a < b {
|
||||
return a
|
||||
}
|
||||
return b
|
||||
}
|
||||
404
internal/app/files/maptile_proxy.go
Normal file
404
internal/app/files/maptile_proxy.go
Normal file
|
|
@ -0,0 +1,404 @@
|
|||
package files
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sort"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"go.uber.org/zap"
|
||||
"golang.org/x/sync/singleflight"
|
||||
)
|
||||
|
||||
// 本文件实现地图缩略图的真实数据源:upload.getWebFile 命中 geo 坐标时,服务端代理
|
||||
// Mapbox Static Images API 抓取一张静态地图并落盘缓存。客户端按 offset/limit 分片下载
|
||||
// 同一文件,字节必须全程一致,因此:
|
||||
// - 抓取成功先原子落盘(temp+rename),所有分片一律从缓存文件读;
|
||||
// - 落盘失败时字节进短 TTL 进程内缓存兜底——分片是顺序请求、singleflight 合并不了,
|
||||
// 没有这层缓存会退化成每分片一次外网抓取,且两次抓取的字节无逐字节一致保证;
|
||||
// - 抓取失败记一个短 TTL 负缓存,期间该 key 直接走确定性占位图,避免同一次下载
|
||||
// 前后分片在「真图/占位图」之间翻转,也避免上游故障时每个分片都打一次外网。
|
||||
//
|
||||
// 资源边界(key 空间由客户端可控的坐标×尺寸组合构成,必须设防):
|
||||
// - 抓取尺寸量化到 32px 档位,收窄 key 空间与缓存基数;
|
||||
// - 磁盘缓存有总量上限,超限按 mtime 从旧到新淘汰(顺带回收崩溃遗留 .tmp);
|
||||
// - 上游抓取有全局速率上限,超限当次退占位图,防止恶意枚举烧 Mapbox 配额。
|
||||
//
|
||||
// 地图不内嵌定位针:TDesktop(historyMapPoint icon)与 DrKLO 都在客户端叠加 marker,
|
||||
// 与官方静态图行为一致。
|
||||
|
||||
const (
|
||||
mapTileFetchTimeout = 15 * time.Second
|
||||
mapTileMaxFetchBytes = 8 << 20 // 防御性上限;640x640@2x PNG 远小于此
|
||||
mapTileFailureTTL = time.Minute
|
||||
mapboxStaticStyleBase = "/styles/v1/mapbox/streets-v12/static"
|
||||
|
||||
// mapTileEdgeStep 是抓取尺寸的量化步长(向上取整);客户端拿到略大的图自适应缩放。
|
||||
mapTileEdgeStep = 32
|
||||
// mapTileMemTTL/mapTileMemMaxBytes 是落盘失败兜底字节缓存的保留期与总量上限。
|
||||
mapTileMemTTL = 10 * time.Minute
|
||||
mapTileMemMaxBytes = 32 << 20
|
||||
// mapTileDiskMaxBytes 是磁盘缓存总量上限;超限按 mtime 淘汰到 90%。
|
||||
mapTileDiskMaxBytes = int64(256 << 20)
|
||||
// mapTileFetchRateLimit 是全局每分钟上游抓取上限(防恶意坐标枚举烧配额)。
|
||||
mapTileFetchRateLimit = 120
|
||||
mapTileFetchRateWindow = time.Minute
|
||||
// mapTileTmpMaxAge 是崩溃遗留 .tmp 的回收阈值。
|
||||
mapTileTmpMaxAge = time.Hour
|
||||
)
|
||||
|
||||
type memTileEntry struct {
|
||||
data []byte
|
||||
at time.Time
|
||||
}
|
||||
|
||||
type mapTileProxy struct {
|
||||
token string
|
||||
baseURL string // 默认 https://api.mapbox.com;测试注入 httptest 地址
|
||||
dir string
|
||||
client *http.Client
|
||||
log *zap.Logger
|
||||
maxDiskBytes int64
|
||||
|
||||
group singleflight.Group
|
||||
sweepMu sync.Mutex
|
||||
|
||||
mu sync.Mutex
|
||||
failures map[string]time.Time // key → 失败时刻(负缓存)
|
||||
memTiles map[string]memTileEntry
|
||||
memBytes int
|
||||
fetchTimes []time.Time // 上游抓取滑动窗口
|
||||
}
|
||||
|
||||
// WithMapboxMapTiles 启用 Mapbox 静态地图代理;token 为空时不启用(保持占位图)。
|
||||
// logger 由 NewService 在全部 Option 应用后统一注入。
|
||||
func WithMapboxMapTiles(token, cacheDir string) Option {
|
||||
return func(s *Service) {
|
||||
if token == "" || cacheDir == "" {
|
||||
return
|
||||
}
|
||||
s.mapTiles = &mapTileProxy{
|
||||
token: token,
|
||||
baseURL: "https://api.mapbox.com",
|
||||
dir: cacheDir,
|
||||
client: &http.Client{Timeout: mapTileFetchTimeout},
|
||||
maxDiskBytes: mapTileDiskMaxBytes,
|
||||
failures: make(map[string]time.Time),
|
||||
memTiles: make(map[string]memTileEntry),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// tile 返回 (lat,long,zoom,w,h,scale) 对应的静态地图字节;入参须已 clamp。
|
||||
func (p *mapTileProxy) tile(lat, long float64, w, h, zoom, scale int) ([]byte, string, error) {
|
||||
// Mapbox 静态图只支持 @2x;scale 3 同样按 2x 抓(客户端按目标尺寸自适应缩放)。
|
||||
retina := ""
|
||||
if scale >= 2 {
|
||||
retina = "@2x"
|
||||
}
|
||||
// 尺寸量化收窄客户端可铸造的 key 空间;客户端按目标矩形自适应缩放略大的图。
|
||||
w = quantizeTileEdge(w)
|
||||
h = quantizeTileEdge(h)
|
||||
key := fmt.Sprintf("v1-%.5f-%.5f-%d-%dx%d%s", lat, long, zoom, w, h, retina)
|
||||
path := p.cachePath(key)
|
||||
if data, err := os.ReadFile(path); err == nil && len(data) > 0 {
|
||||
return data, mapTileMime(data), nil
|
||||
}
|
||||
if data, ok := p.cachedMemTile(key); ok {
|
||||
return data, mapTileMime(data), nil
|
||||
}
|
||||
if p.recentlyFailed(key) {
|
||||
return nil, "", errors.New("map tile fetch in failure backoff")
|
||||
}
|
||||
v, err, _ := p.group.Do(key, func() (any, error) {
|
||||
if data, err := os.ReadFile(path); err == nil && len(data) > 0 {
|
||||
return data, nil
|
||||
}
|
||||
if data, ok := p.cachedMemTile(key); ok {
|
||||
return data, nil
|
||||
}
|
||||
if !p.allowFetch() {
|
||||
// 全局抓取限速:负缓存让该 key 短期稳定走占位图(保持分片字节一致),不打上游。
|
||||
p.markFailed(key)
|
||||
return nil, errors.New("map tile fetch rate limited")
|
||||
}
|
||||
data, err := p.fetch(lat, long, w, h, zoom, retina)
|
||||
if err != nil {
|
||||
p.markFailed(key)
|
||||
return nil, err
|
||||
}
|
||||
if err := p.store(path, data); err != nil {
|
||||
// 分片是顺序请求,singleflight 合并不了后续分片;字节必须进内存缓存兜底,
|
||||
// 否则磁盘持续故障会退化成每分片一次外网抓取且字节无一致性保证。
|
||||
p.rememberMemTile(key, data)
|
||||
p.log.Warn("map tile cache write failed, serving from memory", zap.Error(err), zap.String("key", key))
|
||||
} else {
|
||||
p.sweepDisk()
|
||||
}
|
||||
return data, nil
|
||||
})
|
||||
if err != nil {
|
||||
return nil, "", err
|
||||
}
|
||||
data := v.([]byte)
|
||||
return data, mapTileMime(data), nil
|
||||
}
|
||||
|
||||
// quantizeTileEdge 把边长向上量化到 mapTileEdgeStep 的整数倍(caller 已 clamp 到 16..1024)。
|
||||
func quantizeTileEdge(v int) int {
|
||||
q := ((v + mapTileEdgeStep - 1) / mapTileEdgeStep) * mapTileEdgeStep
|
||||
if q > mapTileMaxEdge {
|
||||
return mapTileMaxEdge
|
||||
}
|
||||
return q
|
||||
}
|
||||
|
||||
// fetch 请求 Mapbox Static Images API。注意 URL 坐标顺序是 {long},{lat}。
|
||||
func (p *mapTileProxy) fetch(lat, long float64, w, h, zoom int, retina string) ([]byte, error) {
|
||||
endpoint := fmt.Sprintf("%s%s/%.5f,%.5f,%d/%dx%d%s?access_token=%s&attribution=false&logo=false",
|
||||
p.baseURL, mapboxStaticStyleBase, long, lat, zoom, w, h, retina, url.QueryEscape(p.token))
|
||||
// 不透传 RPC ctx:singleflight 结果被并发分片共享,单个调用方取消不应拖垮整次抓取。
|
||||
ctx, cancel := context.WithTimeout(context.Background(), mapTileFetchTimeout)
|
||||
defer cancel()
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, endpoint, nil)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("build map tile request: %w", err)
|
||||
}
|
||||
resp, err := p.client.Do(req)
|
||||
if err != nil {
|
||||
// 传输层错误内嵌完整 URL(含 access_token),脱敏后再向上传播/落日志。
|
||||
return nil, fmt.Errorf("fetch map tile: %s", p.redactToken(err.Error()))
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return nil, fmt.Errorf("fetch map tile: upstream status %d", resp.StatusCode)
|
||||
}
|
||||
data, err := io.ReadAll(io.LimitReader(resp.Body, mapTileMaxFetchBytes+1))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("read map tile body: %w", err)
|
||||
}
|
||||
if len(data) == 0 || len(data) > mapTileMaxFetchBytes {
|
||||
return nil, fmt.Errorf("map tile body size invalid: %d", len(data))
|
||||
}
|
||||
if mime := mapTileMime(data); mime == "" {
|
||||
return nil, errors.New("map tile body is not an image")
|
||||
}
|
||||
return data, nil
|
||||
}
|
||||
|
||||
func (p *mapTileProxy) cachePath(key string) string {
|
||||
sum := sha256.Sum256([]byte(key))
|
||||
return filepath.Join(p.dir, hex.EncodeToString(sum[:])+".img")
|
||||
}
|
||||
|
||||
func (p *mapTileProxy) store(path string, data []byte) error {
|
||||
if err := os.MkdirAll(p.dir, 0o755); err != nil {
|
||||
return err
|
||||
}
|
||||
tmp, err := os.CreateTemp(p.dir, "tile-*.tmp")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
tmpName := tmp.Name()
|
||||
if _, err := tmp.Write(data); err != nil {
|
||||
tmp.Close()
|
||||
os.Remove(tmpName)
|
||||
return err
|
||||
}
|
||||
if err := tmp.Close(); err != nil {
|
||||
os.Remove(tmpName)
|
||||
return err
|
||||
}
|
||||
if err := os.Rename(tmpName, path); err != nil {
|
||||
os.Remove(tmpName)
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// redactToken 把错误文本中的 access token(原样或 URL 转义形态)替换为占位符。
|
||||
func (p *mapTileProxy) redactToken(s string) string {
|
||||
if p.token == "" {
|
||||
return s
|
||||
}
|
||||
s = strings.ReplaceAll(s, p.token, "***")
|
||||
if escaped := url.QueryEscape(p.token); escaped != p.token {
|
||||
s = strings.ReplaceAll(s, escaped, "***")
|
||||
}
|
||||
return s
|
||||
}
|
||||
|
||||
// cachedMemTile 返回落盘失败兜底缓存中的字节(TTL 内)。
|
||||
func (p *mapTileProxy) cachedMemTile(key string) ([]byte, bool) {
|
||||
p.mu.Lock()
|
||||
defer p.mu.Unlock()
|
||||
entry, ok := p.memTiles[key]
|
||||
if !ok {
|
||||
return nil, false
|
||||
}
|
||||
if time.Since(entry.at) > mapTileMemTTL {
|
||||
p.memBytes -= len(entry.data)
|
||||
delete(p.memTiles, key)
|
||||
return nil, false
|
||||
}
|
||||
return entry.data, true
|
||||
}
|
||||
|
||||
// rememberMemTile 在落盘失败时暂存字节;超总量按最旧淘汰。
|
||||
func (p *mapTileProxy) rememberMemTile(key string, data []byte) {
|
||||
if len(data) > mapTileMemMaxBytes {
|
||||
return
|
||||
}
|
||||
p.mu.Lock()
|
||||
defer p.mu.Unlock()
|
||||
now := time.Now()
|
||||
for k, entry := range p.memTiles {
|
||||
if now.Sub(entry.at) > mapTileMemTTL {
|
||||
p.memBytes -= len(entry.data)
|
||||
delete(p.memTiles, k)
|
||||
}
|
||||
}
|
||||
if old, ok := p.memTiles[key]; ok {
|
||||
p.memBytes -= len(old.data)
|
||||
}
|
||||
for p.memBytes+len(data) > mapTileMemMaxBytes && len(p.memTiles) > 0 {
|
||||
oldestKey := ""
|
||||
var oldestAt time.Time
|
||||
for k, entry := range p.memTiles {
|
||||
if oldestKey == "" || entry.at.Before(oldestAt) {
|
||||
oldestKey, oldestAt = k, entry.at
|
||||
}
|
||||
}
|
||||
p.memBytes -= len(p.memTiles[oldestKey].data)
|
||||
delete(p.memTiles, oldestKey)
|
||||
}
|
||||
p.memTiles[key] = memTileEntry{data: data, at: now}
|
||||
p.memBytes += len(data)
|
||||
}
|
||||
|
||||
// allowFetch 是全局上游抓取限速(滑动窗口);超限返回 false。
|
||||
func (p *mapTileProxy) allowFetch() bool {
|
||||
p.mu.Lock()
|
||||
defer p.mu.Unlock()
|
||||
now := time.Now()
|
||||
kept := p.fetchTimes[:0]
|
||||
for _, at := range p.fetchTimes {
|
||||
if now.Sub(at) <= mapTileFetchRateWindow {
|
||||
kept = append(kept, at)
|
||||
}
|
||||
}
|
||||
p.fetchTimes = kept
|
||||
if len(p.fetchTimes) >= mapTileFetchRateLimit {
|
||||
return false
|
||||
}
|
||||
p.fetchTimes = append(p.fetchTimes, now)
|
||||
return true
|
||||
}
|
||||
|
||||
// sweepDisk 在新写入后核算缓存目录总量,超限按 mtime 从旧到新淘汰到 90%,
|
||||
// 顺带回收崩溃遗留的过期 .tmp。store 仅发生在上游抓取后(低频),同步扫描可接受。
|
||||
func (p *mapTileProxy) sweepDisk() {
|
||||
p.sweepMu.Lock()
|
||||
defer p.sweepMu.Unlock()
|
||||
entries, err := os.ReadDir(p.dir)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
type tileFile struct {
|
||||
name string
|
||||
size int64
|
||||
mod time.Time
|
||||
}
|
||||
var files []tileFile
|
||||
var total int64
|
||||
now := time.Now()
|
||||
for _, entry := range entries {
|
||||
if entry.IsDir() {
|
||||
continue
|
||||
}
|
||||
info, err := entry.Info()
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
if strings.HasSuffix(entry.Name(), ".tmp") {
|
||||
if now.Sub(info.ModTime()) > mapTileTmpMaxAge {
|
||||
_ = os.Remove(filepath.Join(p.dir, entry.Name()))
|
||||
}
|
||||
continue
|
||||
}
|
||||
if !strings.HasSuffix(entry.Name(), ".img") {
|
||||
continue
|
||||
}
|
||||
files = append(files, tileFile{name: entry.Name(), size: info.Size(), mod: info.ModTime()})
|
||||
total += info.Size()
|
||||
}
|
||||
if total <= p.maxDiskBytes {
|
||||
return
|
||||
}
|
||||
sort.Slice(files, func(i, j int) bool { return files[i].mod.Before(files[j].mod) })
|
||||
target := p.maxDiskBytes * 9 / 10
|
||||
removed := 0
|
||||
for _, f := range files {
|
||||
if total <= target {
|
||||
break
|
||||
}
|
||||
if err := os.Remove(filepath.Join(p.dir, f.name)); err == nil {
|
||||
total -= f.size
|
||||
removed++
|
||||
}
|
||||
}
|
||||
if removed > 0 && p.log != nil {
|
||||
p.log.Info("map tile cache swept", zap.Int("removed", removed), zap.Int64("remaining_bytes", total))
|
||||
}
|
||||
}
|
||||
|
||||
func (p *mapTileProxy) recentlyFailed(key string) bool {
|
||||
p.mu.Lock()
|
||||
defer p.mu.Unlock()
|
||||
at, ok := p.failures[key]
|
||||
if !ok {
|
||||
return false
|
||||
}
|
||||
if time.Since(at) > mapTileFailureTTL {
|
||||
delete(p.failures, key)
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
func (p *mapTileProxy) markFailed(key string) {
|
||||
p.mu.Lock()
|
||||
defer p.mu.Unlock()
|
||||
// 顺手清理过期项,防 map 无界增长(key 空间本身有限:坐标×尺寸枚举)。
|
||||
now := time.Now()
|
||||
for k, at := range p.failures {
|
||||
if now.Sub(at) > mapTileFailureTTL {
|
||||
delete(p.failures, k)
|
||||
}
|
||||
}
|
||||
p.failures[key] = now
|
||||
}
|
||||
|
||||
// mapTileMime 按魔数识别图片类型;非图片返回空串。
|
||||
func mapTileMime(data []byte) string {
|
||||
switch {
|
||||
case len(data) > 8 && data[0] == 0x89 && data[1] == 'P' && data[2] == 'N' && data[3] == 'G':
|
||||
return "image/png"
|
||||
case len(data) > 3 && data[0] == 0xFF && data[1] == 0xD8:
|
||||
return "image/jpeg"
|
||||
case len(data) > 12 && string(data[0:4]) == "RIFF" && string(data[8:12]) == "WEBP":
|
||||
return "image/webp"
|
||||
default:
|
||||
return ""
|
||||
}
|
||||
}
|
||||
243
internal/app/files/maptile_proxy_test.go
Normal file
243
internal/app/files/maptile_proxy_test.go
Normal file
|
|
@ -0,0 +1,243 @@
|
|||
package files
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
// fakePNG 是最小合法 PNG 头 + 填充(只需通过魔数识别,不需要可解码)。
|
||||
var fakePNG = append([]byte{0x89, 'P', 'N', 'G', 0x0D, 0x0A, 0x1A, 0x0A}, bytes.Repeat([]byte{0x42}, 64)...)
|
||||
|
||||
func newProxyTestService(t *testing.T, handler http.Handler) (*Service, *httptest.Server) {
|
||||
t.Helper()
|
||||
srv := httptest.NewServer(handler)
|
||||
t.Cleanup(srv.Close)
|
||||
s := NewService(nil, nil, 2, WithMapboxMapTiles("test-token", t.TempDir()))
|
||||
if s.mapTiles == nil {
|
||||
t.Fatal("map tile proxy not configured")
|
||||
}
|
||||
s.mapTiles.baseURL = srv.URL
|
||||
return s, srv
|
||||
}
|
||||
|
||||
func TestGeoMapTileProxyFetchesAndCaches(t *testing.T) {
|
||||
var calls atomic.Int32
|
||||
var lastPath, lastQuery string
|
||||
s, _ := newProxyTestService(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
calls.Add(1)
|
||||
lastPath = r.URL.Path
|
||||
lastQuery = r.URL.RawQuery
|
||||
w.Header().Set("Content-Type", "image/png")
|
||||
_, _ = w.Write(fakePNG)
|
||||
}))
|
||||
|
||||
first, mime := s.GeoMapTile(39.9042, 116.4074, 256, 128, 15, 2)
|
||||
if mime != "image/png" {
|
||||
t.Fatalf("mime = %q, want image/png", mime)
|
||||
}
|
||||
if !bytes.Equal(first, fakePNG) {
|
||||
t.Fatal("first fetch did not return upstream bytes")
|
||||
}
|
||||
// Mapbox 形态:/styles/v1/mapbox/streets-v12/static/{long},{lat},{zoom}/{w}x{h}@2x
|
||||
if !strings.Contains(lastPath, "/static/116.40740,39.90420,15/256x128@2x") {
|
||||
t.Fatalf("unexpected upstream path: %s", lastPath)
|
||||
}
|
||||
if !strings.Contains(lastQuery, "access_token=test-token") {
|
||||
t.Fatalf("missing access token in query: %s", lastQuery)
|
||||
}
|
||||
|
||||
// 第二次(含分片重复读)必须走磁盘缓存,不再触发上游请求。
|
||||
second, _ := s.GeoMapTile(39.9042, 116.4074, 256, 128, 15, 2)
|
||||
if !bytes.Equal(first, second) {
|
||||
t.Fatal("cached bytes differ from first fetch")
|
||||
}
|
||||
if got := calls.Load(); got != 1 {
|
||||
t.Fatalf("upstream calls = %d, want 1 (second hit must be cached)", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGeoMapTileProxyScaleOneOmitsRetina(t *testing.T) {
|
||||
var lastPath string
|
||||
s, _ := newProxyTestService(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
lastPath = r.URL.Path
|
||||
w.Header().Set("Content-Type", "image/png")
|
||||
_, _ = w.Write(fakePNG)
|
||||
}))
|
||||
s.GeoMapTile(1.5, 2.5, 100, 100, 16, 1)
|
||||
if strings.Contains(lastPath, "@2x") {
|
||||
t.Fatalf("scale=1 must not request retina: %s", lastPath)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGeoMapTileProxyFallsBackToPlaceholder(t *testing.T) {
|
||||
var calls atomic.Int32
|
||||
s, _ := newProxyTestService(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
calls.Add(1)
|
||||
http.Error(w, "boom", http.StatusInternalServerError)
|
||||
}))
|
||||
|
||||
data, mime := s.GeoMapTile(39.9042, 116.4074, 256, 128, 15, 2)
|
||||
if len(data) == 0 || mime != "image/png" {
|
||||
t.Fatalf("fallback placeholder missing: len=%d mime=%q", len(data), mime)
|
||||
}
|
||||
// 占位图确定性:回退路径必须与纯占位服务字节一致(分片续传一致性)。
|
||||
plain := NewService(nil, nil, 2)
|
||||
expected, _ := plain.GeoMapTile(39.9042, 116.4074, 256, 128, 15, 2)
|
||||
if !bytes.Equal(data, expected) {
|
||||
t.Fatal("fallback placeholder differs from deterministic rendering")
|
||||
}
|
||||
|
||||
// 负缓存:失败后的后续分片请求不应继续打上游。
|
||||
again, _ := s.GeoMapTile(39.9042, 116.4074, 256, 128, 15, 2)
|
||||
if !bytes.Equal(again, expected) {
|
||||
t.Fatal("placeholder must stay byte-identical during failure backoff")
|
||||
}
|
||||
if got := calls.Load(); got != 1 {
|
||||
t.Fatalf("upstream calls = %d, want 1 (failure must be negative-cached)", got)
|
||||
}
|
||||
|
||||
// 负缓存过期后允许重试并恢复真实地图。
|
||||
s.mapTiles.mu.Lock()
|
||||
for k := range s.mapTiles.failures {
|
||||
s.mapTiles.failures[k] = time.Now().Add(-2 * mapTileFailureTTL)
|
||||
}
|
||||
s.mapTiles.mu.Unlock()
|
||||
// 上游恢复。
|
||||
s.mapTiles.baseURL = newRecoveredUpstream(t)
|
||||
recovered, _ := s.GeoMapTile(39.9042, 116.4074, 256, 128, 15, 2)
|
||||
if !bytes.Equal(recovered, fakePNG) {
|
||||
t.Fatal("proxy must recover after failure TTL expires")
|
||||
}
|
||||
}
|
||||
|
||||
func newRecoveredUpstream(t *testing.T) string {
|
||||
t.Helper()
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "image/png")
|
||||
_, _ = w.Write(fakePNG)
|
||||
}))
|
||||
t.Cleanup(srv.Close)
|
||||
return srv.URL
|
||||
}
|
||||
|
||||
// 落盘失败时字节必须进内存兜底缓存:顺序分片不再逐片打上游,且字节全程一致。
|
||||
func TestGeoMapTileProxyStoreFailureServesFromMemory(t *testing.T) {
|
||||
var calls atomic.Int32
|
||||
s, _ := newProxyTestService(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
calls.Add(1)
|
||||
w.Header().Set("Content-Type", "image/png")
|
||||
_, _ = w.Write(fakePNG)
|
||||
}))
|
||||
// 让缓存目录路径指向一个普通文件 → MkdirAll/写盘必然失败。
|
||||
blocked := filepath.Join(t.TempDir(), "not-a-dir")
|
||||
if err := os.WriteFile(blocked, []byte("x"), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
s.mapTiles.dir = blocked
|
||||
|
||||
first, _ := s.GeoMapTile(39.9042, 116.4074, 256, 128, 15, 2)
|
||||
second, _ := s.GeoMapTile(39.9042, 116.4074, 256, 128, 15, 2)
|
||||
if !bytes.Equal(first, fakePNG) || !bytes.Equal(second, fakePNG) {
|
||||
t.Fatal("store-failure path must keep serving upstream bytes")
|
||||
}
|
||||
if got := calls.Load(); got != 1 {
|
||||
t.Fatalf("upstream calls = %d, want 1 (memory cache must absorb后续分片)", got)
|
||||
}
|
||||
}
|
||||
|
||||
// 抓取尺寸量化到 32px 档位,收窄客户端可铸造的缓存 key 空间。
|
||||
func TestGeoMapTileProxyQuantizesFetchDimensions(t *testing.T) {
|
||||
var mu sync.Mutex
|
||||
var lastPath string
|
||||
s, _ := newProxyTestService(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
mu.Lock()
|
||||
lastPath = r.URL.Path
|
||||
mu.Unlock()
|
||||
w.Header().Set("Content-Type", "image/png")
|
||||
_, _ = w.Write(fakePNG)
|
||||
}))
|
||||
s.GeoMapTile(1.5, 2.5, 100, 50, 15, 2)
|
||||
mu.Lock()
|
||||
defer mu.Unlock()
|
||||
if !strings.Contains(lastPath, "/128x64@2x") {
|
||||
t.Fatalf("dimensions not quantized to 32px steps: %s", lastPath)
|
||||
}
|
||||
}
|
||||
|
||||
// 磁盘缓存超总量上限后按 mtime 从旧到新淘汰。
|
||||
func TestGeoMapTileProxyDiskSweepEvictsOldest(t *testing.T) {
|
||||
s, _ := newProxyTestService(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "image/png")
|
||||
_, _ = w.Write(fakePNG)
|
||||
}))
|
||||
s.mapTiles.maxDiskBytes = int64(len(fakePNG)*2 + 8) // 容得下 2 张,第 3 张触发淘汰
|
||||
for i := 0; i < 3; i++ {
|
||||
s.GeoMapTile(10+float64(i), 20, 128, 128, 15, 1)
|
||||
time.Sleep(20 * time.Millisecond) // 保证 mtime 可区分
|
||||
}
|
||||
entries, err := os.ReadDir(s.mapTiles.dir)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
var total int64
|
||||
count := 0
|
||||
for _, entry := range entries {
|
||||
info, err := entry.Info()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
total += info.Size()
|
||||
count++
|
||||
}
|
||||
if count >= 3 || total > s.mapTiles.maxDiskBytes {
|
||||
t.Fatalf("sweep did not evict: files=%d total=%d max=%d", count, total, s.mapTiles.maxDiskBytes)
|
||||
}
|
||||
}
|
||||
|
||||
// 全局抓取限速:超限的新 key 退占位图且不打上游。
|
||||
func TestGeoMapTileProxyFetchRateLimit(t *testing.T) {
|
||||
var calls atomic.Int32
|
||||
s, _ := newProxyTestService(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
calls.Add(1)
|
||||
w.Header().Set("Content-Type", "image/png")
|
||||
_, _ = w.Write(fakePNG)
|
||||
}))
|
||||
// 预填满滑动窗口。
|
||||
s.mapTiles.mu.Lock()
|
||||
now := time.Now()
|
||||
for i := 0; i < mapTileFetchRateLimit; i++ {
|
||||
s.mapTiles.fetchTimes = append(s.mapTiles.fetchTimes, now)
|
||||
}
|
||||
s.mapTiles.mu.Unlock()
|
||||
|
||||
data, mime := s.GeoMapTile(33.3, 44.4, 128, 128, 15, 1)
|
||||
plain := NewService(nil, nil, 2)
|
||||
expected, _ := plain.GeoMapTile(33.3, 44.4, 128, 128, 15, 1)
|
||||
if !bytes.Equal(data, expected) || mime != "image/png" {
|
||||
t.Fatal("rate-limited request must fall back to deterministic placeholder")
|
||||
}
|
||||
if got := calls.Load(); got != 0 {
|
||||
t.Fatalf("upstream calls = %d, want 0 when rate limited", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGeoMapTileProxyRejectsNonImageBody(t *testing.T) {
|
||||
s, _ := newProxyTestService(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "text/html")
|
||||
_, _ = w.Write([]byte("<html>not a map</html>"))
|
||||
}))
|
||||
data, mime := s.GeoMapTile(10, 20, 128, 128, 15, 1)
|
||||
plain := NewService(nil, nil, 2)
|
||||
expected, _ := plain.GeoMapTile(10, 20, 128, 128, 15, 1)
|
||||
if !bytes.Equal(data, expected) || mime != "image/png" {
|
||||
t.Fatal("non-image upstream body must fall back to placeholder")
|
||||
}
|
||||
}
|
||||
48
internal/app/files/maptile_test.go
Normal file
48
internal/app/files/maptile_test.go
Normal file
|
|
@ -0,0 +1,48 @@
|
|||
package files
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"image/png"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestGeoMapTileDeterministicPNG(t *testing.T) {
|
||||
var s *Service
|
||||
first, mime := s.GeoMapTile(39.9042, 116.4074, 256, 128, 15, 2)
|
||||
second, _ := s.GeoMapTile(39.9042, 116.4074, 256, 128, 15, 2)
|
||||
if mime != "image/png" {
|
||||
t.Fatalf("mime = %q, want image/png", mime)
|
||||
}
|
||||
if !bytes.Equal(first, second) {
|
||||
t.Fatal("map tile must be byte-identical for identical input (chunked download consistency)")
|
||||
}
|
||||
img, err := png.Decode(bytes.NewReader(first))
|
||||
if err != nil {
|
||||
t.Fatalf("decode png: %v", err)
|
||||
}
|
||||
if img.Bounds().Dx() != 512 || img.Bounds().Dy() != 256 {
|
||||
t.Fatalf("tile dims = %dx%d, want 512x256 (w*scale x h*scale)", img.Bounds().Dx(), img.Bounds().Dy())
|
||||
}
|
||||
}
|
||||
|
||||
func TestGeoMapTileClampsBounds(t *testing.T) {
|
||||
var s *Service
|
||||
tile, _ := s.GeoMapTile(0, 0, 99999, -5, 99, 9)
|
||||
img, err := png.Decode(bytes.NewReader(tile))
|
||||
if err != nil {
|
||||
t.Fatalf("decode png: %v", err)
|
||||
}
|
||||
// w clamp 1024、h clamp 16、scale clamp 3。
|
||||
if img.Bounds().Dx() != 1024*3 || img.Bounds().Dy() != 16*3 {
|
||||
t.Fatalf("tile dims = %dx%d, want 3072x48", img.Bounds().Dx(), img.Bounds().Dy())
|
||||
}
|
||||
}
|
||||
|
||||
func TestGeoMapTileDiffersByLocation(t *testing.T) {
|
||||
var s *Service
|
||||
a, _ := s.GeoMapTile(39.9042, 116.4074, 128, 128, 15, 1)
|
||||
b, _ := s.GeoMapTile(31.2304, 121.4737, 128, 128, 15, 1)
|
||||
if bytes.Equal(a, b) {
|
||||
t.Fatal("different locations should render different tiles")
|
||||
}
|
||||
}
|
||||
243
internal/app/files/mp4faststart.go
Normal file
243
internal/app/files/mp4faststart.go
Normal file
|
|
@ -0,0 +1,243 @@
|
|||
package files
|
||||
|
||||
import "encoding/binary"
|
||||
|
||||
// faststartMP4 把 MP4 的 moov 原子移到 mdat 之前(faststart),并把 moov 内所有
|
||||
// stco/co64 chunk 偏移整体加上 moov 的大小(因为 moov 插到 mdat 前会把媒体数据整体后移)。
|
||||
// 返回 (newData, changed)。非 MP4 / 已 faststart / 任何解析异常时返回 (data, false) 原样,
|
||||
// 绝不破坏数据——这是上传热路径,宁可不优化也不能把好视频改坏。
|
||||
//
|
||||
// 背景:TDesktop 的 story 流式播放路径无法处理 moov 在文件末尾的视频(av_read_frame
|
||||
// 报 Invalid data),普通 Telegram 客户端上传前会 faststart。telesrv 在落盘前补这一步,
|
||||
// 让 supports_streaming=true 的承诺对所有客户端成立。不转码、保留原编码(含 HEVC)。
|
||||
func faststartMP4(data []byte) ([]byte, bool) {
|
||||
boxes, ok := parseTopLevelBoxes(data)
|
||||
if !ok {
|
||||
return data, false
|
||||
}
|
||||
ftypIdx, moovIdx, firstMdatIdx := -1, -1, -1
|
||||
for i, b := range boxes {
|
||||
switch b.typ {
|
||||
case "ftyp":
|
||||
if ftypIdx < 0 {
|
||||
ftypIdx = i
|
||||
}
|
||||
case "moov":
|
||||
if moovIdx < 0 {
|
||||
moovIdx = i
|
||||
}
|
||||
case "mdat":
|
||||
if firstMdatIdx < 0 {
|
||||
firstMdatIdx = i
|
||||
}
|
||||
}
|
||||
}
|
||||
// 必须有 ftyp(且在最前)、moov、mdat;moov 已在 mdat 前则已 faststart。
|
||||
if ftypIdx != 0 || moovIdx < 0 || firstMdatIdx < 0 {
|
||||
return data, false
|
||||
}
|
||||
if moovIdx < firstMdatIdx {
|
||||
return data, false
|
||||
}
|
||||
|
||||
// 拷出 moov(独立底层数组,后续就地改偏移不影响原 data)。
|
||||
moov := append([]byte(nil), data[boxes[moovIdx].start:boxes[moovIdx].end]...)
|
||||
moovSize := int64(len(moov))
|
||||
if !patchChunkOffsets(moov, moovSize) {
|
||||
return data, false
|
||||
}
|
||||
|
||||
// 重组:ftyp + moov + 其余 box(除 ftyp/moov)按原顺序。
|
||||
out := make([]byte, 0, len(data))
|
||||
out = append(out, data[boxes[ftypIdx].start:boxes[ftypIdx].end]...)
|
||||
out = append(out, moov...)
|
||||
for i, b := range boxes {
|
||||
if i == ftypIdx || i == moovIdx {
|
||||
continue
|
||||
}
|
||||
out = append(out, data[b.start:b.end]...)
|
||||
}
|
||||
if len(out) != len(data) {
|
||||
// 长度必须守恒(只搬不改大小);不守恒说明哪里算错,保守放弃。
|
||||
return data, false
|
||||
}
|
||||
return out, true
|
||||
}
|
||||
|
||||
// mp4Layout 是只读 box 头得出的顶层结构信息,用于在不读取整段媒体的前提下判断是否需要
|
||||
// faststart 以及如何流式重写。
|
||||
type mp4Layout struct {
|
||||
needsFaststart bool // moov 在 mdat 之后
|
||||
moovIsLast bool // moov 是最后一个顶层 box(可走流式重写)
|
||||
ftypStart, ftypEnd int64
|
||||
moovStart, moovEnd int64
|
||||
}
|
||||
|
||||
// inspectMP4Layout 只读顶层 box 头(每个 ≤16 字节,跳过 box 负载)来判定结构,避免为了
|
||||
// 「检查是否已 faststart」而把整段视频读进内存。readAt(off, n) 读取 [off, off+n) 字节。
|
||||
// 非 MP4 / 结构异常返回 (·, false)。
|
||||
func inspectMP4Layout(size int64, readAt func(off, n int64) ([]byte, error)) (mp4Layout, bool) {
|
||||
l := mp4Layout{moovStart: -1, moovEnd: -1}
|
||||
mdatStart := int64(-1)
|
||||
ftypSeen := false
|
||||
p := int64(0)
|
||||
for boxes := 0; p+8 <= size; boxes++ {
|
||||
if boxes > 1024 { // 顶层 box 数量上限,防御异常文件
|
||||
return mp4Layout{}, false
|
||||
}
|
||||
hdr, err := readAt(p, 16)
|
||||
if err != nil || int64(len(hdr)) < 8 {
|
||||
return mp4Layout{}, false
|
||||
}
|
||||
boxSize := int64(binary.BigEndian.Uint32(hdr[0:4]))
|
||||
typ := string(hdr[4:8])
|
||||
switch {
|
||||
case boxSize == 1:
|
||||
if len(hdr) < 16 {
|
||||
return mp4Layout{}, false
|
||||
}
|
||||
boxSize = int64(binary.BigEndian.Uint64(hdr[8:16]))
|
||||
case boxSize == 0:
|
||||
boxSize = size - p
|
||||
}
|
||||
if boxSize < 8 || p+boxSize > size {
|
||||
return mp4Layout{}, false
|
||||
}
|
||||
switch typ {
|
||||
case "ftyp":
|
||||
if p != 0 {
|
||||
return mp4Layout{}, false // ftyp 必须在最前
|
||||
}
|
||||
ftypSeen = true
|
||||
l.ftypStart, l.ftypEnd = p, p+boxSize
|
||||
case "moov":
|
||||
if l.moovStart < 0 {
|
||||
l.moovStart, l.moovEnd = p, p+boxSize
|
||||
}
|
||||
case "mdat":
|
||||
if mdatStart < 0 {
|
||||
mdatStart = p
|
||||
}
|
||||
}
|
||||
p += boxSize
|
||||
}
|
||||
if p != size || !ftypSeen || l.moovStart < 0 || mdatStart < 0 {
|
||||
return mp4Layout{}, false
|
||||
}
|
||||
l.needsFaststart = l.moovStart > mdatStart
|
||||
l.moovIsLast = l.moovEnd == size
|
||||
return l, true
|
||||
}
|
||||
|
||||
type boxRef struct {
|
||||
typ string
|
||||
start, end int
|
||||
}
|
||||
|
||||
// parseTopLevelBoxes 顺序解析顶层 box,要求恰好无缝覆盖整个 data,否则视为非法不处理。
|
||||
func parseTopLevelBoxes(data []byte) ([]boxRef, bool) {
|
||||
var boxes []boxRef
|
||||
p := 0
|
||||
for p+8 <= len(data) {
|
||||
size := int(binary.BigEndian.Uint32(data[p : p+4]))
|
||||
typ := string(data[p+4 : p+8])
|
||||
switch {
|
||||
case size == 1:
|
||||
if p+16 > len(data) {
|
||||
return nil, false
|
||||
}
|
||||
size64 := binary.BigEndian.Uint64(data[p+8 : p+16])
|
||||
size = int(size64)
|
||||
case size == 0:
|
||||
size = len(data) - p
|
||||
}
|
||||
if size < 8 || p+size > len(data) {
|
||||
return nil, false
|
||||
}
|
||||
boxes = append(boxes, boxRef{typ: typ, start: p, end: p + size})
|
||||
p += size
|
||||
}
|
||||
if p != len(data) {
|
||||
return nil, false
|
||||
}
|
||||
return boxes, true
|
||||
}
|
||||
|
||||
// patchChunkOffsets 递归进入 moov 的容器 box,把 stco/co64 的每个偏移 += delta。
|
||||
func patchChunkOffsets(box []byte, delta int64) bool {
|
||||
if len(box) < 8 {
|
||||
return false
|
||||
}
|
||||
p := 8 // 跳过自身 box 头(moov 用 32 位 size,极少 64 位;若 64 位则下面 walk 仍从 8 起会错→由调用方 moov 头守恒保证)
|
||||
for p+8 <= len(box) {
|
||||
size := int(binary.BigEndian.Uint32(box[p : p+4]))
|
||||
typ := string(box[p+4 : p+8])
|
||||
hdr := 8
|
||||
switch {
|
||||
case size == 1:
|
||||
if p+16 > len(box) {
|
||||
return false
|
||||
}
|
||||
size = int(binary.BigEndian.Uint64(box[p+8 : p+16]))
|
||||
hdr = 16
|
||||
case size == 0:
|
||||
size = len(box) - p
|
||||
}
|
||||
if size < hdr || p+size > len(box) {
|
||||
return false
|
||||
}
|
||||
child := box[p : p+size]
|
||||
switch typ {
|
||||
case "stco":
|
||||
if !patchStco(child, delta) {
|
||||
return false
|
||||
}
|
||||
case "co64":
|
||||
if !patchCo64(child, delta) {
|
||||
return false
|
||||
}
|
||||
case "trak", "mdia", "minf", "stbl", "edts":
|
||||
if !patchChunkOffsets(child, delta) {
|
||||
return false
|
||||
}
|
||||
}
|
||||
p += size
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
// patchStco:stco = [size(4)][type(4)][version+flags(4)][entry_count(4)][offsets 4*count]。
|
||||
func patchStco(box []byte, delta int64) bool {
|
||||
if len(box) < 16 {
|
||||
return false
|
||||
}
|
||||
count := binary.BigEndian.Uint32(box[12:16])
|
||||
off := 16
|
||||
if int64(off)+int64(count)*4 > int64(len(box)) {
|
||||
return false
|
||||
}
|
||||
for i := uint32(0); i < count; i++ {
|
||||
v := binary.BigEndian.Uint32(box[off : off+4])
|
||||
binary.BigEndian.PutUint32(box[off:off+4], uint32(int64(v)+delta))
|
||||
off += 4
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
// patchCo64:co64 entries 为 8 字节。
|
||||
func patchCo64(box []byte, delta int64) bool {
|
||||
if len(box) < 16 {
|
||||
return false
|
||||
}
|
||||
count := binary.BigEndian.Uint32(box[12:16])
|
||||
off := 16
|
||||
if int64(off)+int64(count)*8 > int64(len(box)) {
|
||||
return false
|
||||
}
|
||||
for i := uint32(0); i < count; i++ {
|
||||
v := binary.BigEndian.Uint64(box[off : off+8])
|
||||
binary.BigEndian.PutUint64(box[off:off+8], v+uint64(delta))
|
||||
off += 8
|
||||
}
|
||||
return true
|
||||
}
|
||||
205
internal/app/files/mp4faststart_test.go
Normal file
205
internal/app/files/mp4faststart_test.go
Normal file
|
|
@ -0,0 +1,205 @@
|
|||
package files
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/binary"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// box 构造一个 MP4 box:[size(4)][type(4)][payload]。
|
||||
func box(typ string, payload []byte) []byte {
|
||||
b := make([]byte, 8+len(payload))
|
||||
binary.BigEndian.PutUint32(b[0:4], uint32(8+len(payload)))
|
||||
copy(b[4:8], typ)
|
||||
copy(b[8:], payload)
|
||||
return b
|
||||
}
|
||||
|
||||
// stcoBox 构造一个含单个 chunk 偏移的 stco:[ver+flags(4)][count(4)=1][offset(4)]。
|
||||
func stcoBox(offset uint32) []byte {
|
||||
p := make([]byte, 12)
|
||||
binary.BigEndian.PutUint32(p[4:8], 1) // entry_count
|
||||
binary.BigEndian.PutUint32(p[8:12], offset)
|
||||
return box("stco", p)
|
||||
}
|
||||
|
||||
// buildMoovEndMP4 构造一个 moov 在末尾、stco 指向 mdat 内某偏移的最小合法 MP4。
|
||||
// 返回 (mp4, mdatPayloadAbsOffset)。
|
||||
func buildMoovEndMP4(mdatPayload []byte) ([]byte, uint32) {
|
||||
ftyp := box("ftyp", []byte("isom\x00\x00\x02\x00"))
|
||||
mdat := box("mdat", mdatPayload)
|
||||
mdatAbs := uint32(len(ftyp) + 8) // mdat payload 紧跟 mdat 头(8 字节)
|
||||
// moov→trak→mdia→minf→stbl→stco(offset=mdatAbs)
|
||||
stbl := box("stbl", stcoBox(mdatAbs))
|
||||
minf := box("minf", stbl)
|
||||
mdia := box("mdia", minf)
|
||||
trak := box("trak", mdia)
|
||||
moov := box("moov", trak)
|
||||
out := append([]byte(nil), ftyp...)
|
||||
out = append(out, mdat...)
|
||||
out = append(out, moov...)
|
||||
return out, mdatAbs
|
||||
}
|
||||
|
||||
func TestFaststartMP4MovesMoovAndFixesOffsets(t *testing.T) {
|
||||
marker := []byte("THE-REAL-CHUNK-DATA-HERE")
|
||||
// mdat payload:前面填充 + marker,stco 指向 marker 的绝对偏移。
|
||||
pad := bytes.Repeat([]byte{0xAB}, 40)
|
||||
mdatPayload := append(append([]byte(nil), pad...), marker...)
|
||||
_, mdatAbs := buildMoovEndMP4(mdatPayload)
|
||||
markerAbs := mdatAbs + uint32(len(pad)) // marker 在原文件里的绝对偏移
|
||||
|
||||
// 原 stco 指向 mdatAbs(mdat payload 头)。这里把 stco 改成指向 marker 以便断言。
|
||||
in2, _ := buildMoovEndMP4Marker(mdatPayload, markerAbs)
|
||||
// 校验前置:原文件里 markerAbs 处确实是 marker。
|
||||
if !bytes.Equal(in2[markerAbs:markerAbs+uint32(len(marker))], marker) {
|
||||
t.Fatalf("setup: 原文件 markerAbs 处非 marker")
|
||||
}
|
||||
|
||||
out, changed := faststartMP4(in2)
|
||||
if !changed {
|
||||
t.Fatalf("changed = false, want true(moov 在末尾应被搬动)")
|
||||
}
|
||||
if len(out) != len(in2) {
|
||||
t.Fatalf("size 不守恒: in=%d out=%d", len(in2), len(out))
|
||||
}
|
||||
// 输出顺序:ftyp, moov, mdat。
|
||||
boxes, ok := parseTopLevelBoxes(out)
|
||||
if !ok || len(boxes) != 3 || boxes[0].typ != "ftyp" || boxes[1].typ != "moov" || boxes[2].typ != "mdat" {
|
||||
t.Fatalf("输出顶层顺序错: %+v ok=%v", boxes, ok)
|
||||
}
|
||||
// 取出输出 moov 里的 stco 偏移,应指向输出文件里仍是 marker 的位置。
|
||||
newOff := readSingleStcoOffset(t, out[boxes[1].start:boxes[1].end])
|
||||
if int(newOff)+len(marker) > len(out) {
|
||||
t.Fatalf("新偏移越界: %d", newOff)
|
||||
}
|
||||
got := out[newOff : int(newOff)+len(marker)]
|
||||
if !bytes.Equal(got, marker) {
|
||||
t.Fatalf("新 stco 偏移 %d 指向 %q, want %q(偏移修正错误)", newOff, got, marker)
|
||||
}
|
||||
// 新偏移 = 原偏移 + moovSize。
|
||||
moovSize := boxes[1].end - boxes[1].start
|
||||
if int(newOff) != int(markerAbs)+moovSize {
|
||||
t.Fatalf("新偏移 = %d, want 原 %d + moovSize %d", newOff, markerAbs, moovSize)
|
||||
}
|
||||
}
|
||||
|
||||
func TestFaststartMP4NoopWhenAlreadyFaststart(t *testing.T) {
|
||||
// moov 在 mdat 前 → 已 faststart,应原样返回。
|
||||
ftyp := box("ftyp", []byte("isom\x00\x00\x02\x00"))
|
||||
stbl := box("stbl", stcoBox(100))
|
||||
moov := box("moov", box("trak", box("mdia", box("minf", stbl))))
|
||||
mdat := box("mdat", bytes.Repeat([]byte{1}, 64))
|
||||
in := append(append(append([]byte(nil), ftyp...), moov...), mdat...)
|
||||
out, changed := faststartMP4(in)
|
||||
if changed {
|
||||
t.Fatalf("已 faststart 不应改动")
|
||||
}
|
||||
if !bytes.Equal(out, in) {
|
||||
t.Fatalf("应原样返回")
|
||||
}
|
||||
}
|
||||
|
||||
func TestInspectMP4LayoutDetectsMoovAtEnd(t *testing.T) {
|
||||
in, _ := buildMoovEndMP4(bytes.Repeat([]byte{0xCD}, 100))
|
||||
readAt := func(off, n int64) ([]byte, error) {
|
||||
if off+n > int64(len(in)) {
|
||||
n = int64(len(in)) - off
|
||||
}
|
||||
return in[off : off+n], nil
|
||||
}
|
||||
l, ok := inspectMP4Layout(int64(len(in)), readAt)
|
||||
if !ok {
|
||||
t.Fatalf("inspect 失败")
|
||||
}
|
||||
if !l.needsFaststart {
|
||||
t.Fatalf("moov 在末尾应 needsFaststart=true")
|
||||
}
|
||||
if !l.moovIsLast {
|
||||
t.Fatalf("moov 是最后一个 box 应 moovIsLast=true")
|
||||
}
|
||||
if l.ftypStart != 0 || l.moovEnd != int64(len(in)) {
|
||||
t.Fatalf("range 不对: %+v (len=%d)", l, len(in))
|
||||
}
|
||||
}
|
||||
|
||||
func TestInspectMP4LayoutNoFaststartWhenMoovFirst(t *testing.T) {
|
||||
ftyp := box("ftyp", []byte("isom\x00\x00\x02\x00"))
|
||||
moov := box("moov", box("trak", box("mdia", box("minf", box("stbl", stcoBox(100))))))
|
||||
mdat := box("mdat", bytes.Repeat([]byte{1}, 100))
|
||||
in := append(append(append([]byte(nil), ftyp...), moov...), mdat...)
|
||||
readAt := func(off, n int64) ([]byte, error) {
|
||||
if off+n > int64(len(in)) {
|
||||
n = int64(len(in)) - off
|
||||
}
|
||||
return in[off : off+n], nil
|
||||
}
|
||||
l, ok := inspectMP4Layout(int64(len(in)), readAt)
|
||||
if !ok || l.needsFaststart {
|
||||
t.Fatalf("moov 在前应 needsFaststart=false, got ok=%v %+v", ok, l)
|
||||
}
|
||||
}
|
||||
|
||||
// 流式拼接(ftyp + patched moov + 中段)必须与全量 faststartMP4 输出逐字节一致。
|
||||
func TestStreamingAssemblyMatchesFullRewrite(t *testing.T) {
|
||||
in, _ := buildMoovEndMP4(bytes.Repeat([]byte{0xEE}, 256))
|
||||
readAt := func(off, n int64) ([]byte, error) { return in[off : off+n], nil }
|
||||
l, ok := inspectMP4Layout(int64(len(in)), readAt)
|
||||
if !ok || !l.needsFaststart || !l.moovIsLast {
|
||||
t.Fatalf("setup: %+v ok=%v", l, ok)
|
||||
}
|
||||
// 流式版本:ftyp + patched moov + in[ftypEnd:moovStart]
|
||||
ftyp := append([]byte(nil), in[l.ftypStart:l.ftypEnd]...)
|
||||
moov := append([]byte(nil), in[l.moovStart:l.moovEnd]...)
|
||||
if !patchChunkOffsets(moov, l.moovEnd-l.moovStart) {
|
||||
t.Fatalf("patch 失败")
|
||||
}
|
||||
var streaming []byte
|
||||
streaming = append(streaming, ftyp...)
|
||||
streaming = append(streaming, moov...)
|
||||
streaming = append(streaming, in[l.ftypEnd:l.moovStart]...)
|
||||
|
||||
full, changed := faststartMP4(in)
|
||||
if !changed {
|
||||
t.Fatalf("full 应 changed")
|
||||
}
|
||||
if !bytes.Equal(streaming, full) {
|
||||
t.Fatalf("流式拼接与全量重写不一致: streaming=%d full=%d", len(streaming), len(full))
|
||||
}
|
||||
}
|
||||
|
||||
func TestFaststartMP4NoopOnNonMP4(t *testing.T) {
|
||||
for _, data := range [][]byte{
|
||||
nil,
|
||||
[]byte("not an mp4 at all"),
|
||||
{0, 0, 0, 4}, // size 太小
|
||||
} {
|
||||
if out, changed := faststartMP4(data); changed || !bytes.Equal(out, data) {
|
||||
t.Fatalf("非 MP4 应原样不动: changed=%v", changed)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// buildMoovEndMP4Marker 同 buildMoovEndMP4,但 stco 指向给定绝对偏移。
|
||||
func buildMoovEndMP4Marker(mdatPayload []byte, stcoOffset uint32) ([]byte, uint32) {
|
||||
ftyp := box("ftyp", []byte("isom\x00\x00\x02\x00"))
|
||||
mdat := box("mdat", mdatPayload)
|
||||
mdatAbs := uint32(len(ftyp) + 8)
|
||||
stbl := box("stbl", stcoBox(stcoOffset))
|
||||
moov := box("moov", box("trak", box("mdia", box("minf", stbl))))
|
||||
out := append([]byte(nil), ftyp...)
|
||||
out = append(out, mdat...)
|
||||
out = append(out, moov...)
|
||||
return out, mdatAbs
|
||||
}
|
||||
|
||||
func readSingleStcoOffset(t *testing.T, moov []byte) uint32 {
|
||||
t.Helper()
|
||||
idx := bytes.Index(moov, []byte("stco"))
|
||||
if idx < 0 {
|
||||
t.Fatalf("moov 里找不到 stco")
|
||||
}
|
||||
// stco 头后:type(4) 已在 idx;payload 从 idx+4 起 = ver+flags(4)+count(4)+offset(4)
|
||||
off := idx + 4 + 4 + 4
|
||||
return binary.BigEndian.Uint32(moov[off : off+4])
|
||||
}
|
||||
|
|
@ -7,11 +7,20 @@ import (
|
|||
"encoding/binary"
|
||||
"fmt"
|
||||
"image"
|
||||
"image/color"
|
||||
"io"
|
||||
stddraw "image/draw"
|
||||
_ "image/jpeg" // 注册 jpeg DecodeConfig,用于读取上传头像/图片尺寸
|
||||
_ "image/png" // 注册 png DecodeConfig
|
||||
"image/png"
|
||||
"math"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"telesrv/internal/domain"
|
||||
|
||||
"go.uber.org/zap"
|
||||
xdraw "golang.org/x/image/draw"
|
||||
_ "golang.org/x/image/webp" // 注册 webp Decode,用于 custom emoji / sticker 静态缩略图合成
|
||||
)
|
||||
|
||||
// 头像与图片消息共用的尺寸 type:'a' 小图(≤160),'c' 大图,'x' 通用下载尺寸。
|
||||
|
|
@ -24,17 +33,10 @@ func (s *Service) UploadProfilePhoto(ctx context.Context, ownerType domain.PeerT
|
|||
|
||||
// UploadProfilePhotoKind stores a profile or fallback photo and makes it current for that kind.
|
||||
func (s *Service) UploadProfilePhotoKind(ctx context.Context, ownerType domain.PeerType, ownerID int64, kind domain.ProfilePhotoKind, file domain.UploadedFileRef, date int) (domain.Photo, error) {
|
||||
data, err := s.assembleUpload(ctx, file.OwnerUserID, file.FileID, file.Parts)
|
||||
if err != nil {
|
||||
return domain.Photo{}, err
|
||||
}
|
||||
if len(data) == 0 {
|
||||
return domain.Photo{}, domain.ErrPhotoInvalid
|
||||
}
|
||||
if date == 0 {
|
||||
date = int(time.Now().Unix())
|
||||
}
|
||||
photo, err := s.createPhoto(ctx, data, photoSizeSpecsForAvatar(data))
|
||||
photo, err := s.CreateAvatarFromUpload(ctx, file)
|
||||
if err != nil {
|
||||
return domain.Photo{}, err
|
||||
}
|
||||
|
|
@ -56,6 +58,14 @@ func (s *Service) CreatePhotoFromUpload(ctx context.Context, file domain.Uploade
|
|||
return s.createPhoto(ctx, data, photoSizeSpecsForMessage(data))
|
||||
}
|
||||
|
||||
// CreatePhotoFromBytes stores already-fetched image bytes as a message Photo.
|
||||
func (s *Service) CreatePhotoFromBytes(ctx context.Context, data []byte) (domain.Photo, error) {
|
||||
if len(data) == 0 {
|
||||
return domain.Photo{}, domain.ErrPhotoInvalid
|
||||
}
|
||||
return s.createPhoto(ctx, data, photoSizeSpecsForMessage(data))
|
||||
}
|
||||
|
||||
// GetPhoto 按 id 返回已存储照片。
|
||||
func (s *Service) GetPhoto(ctx context.Context, id int64) (domain.Photo, bool, error) {
|
||||
return s.media.GetPhoto(ctx, id)
|
||||
|
|
@ -79,25 +89,137 @@ func (s *Service) CreateAvatarFromUpload(ctx context.Context, file domain.Upload
|
|||
return s.createPhoto(ctx, data, photoSizeSpecsForAvatar(data))
|
||||
}
|
||||
|
||||
// CreateAvatarVideoFromUpload stores an animated profile video as photo.video_sizes.
|
||||
func (s *Service) CreateAvatarVideoFromUpload(ctx context.Context, file domain.UploadedFileRef, videoStartTs float64) (domain.Photo, error) {
|
||||
return s.createAvatarVideoFromUpload(ctx, file, videoStartTs, nil)
|
||||
}
|
||||
|
||||
// CreateAvatarVideoMarkupFromUpload stores Android-style generated avatar video plus its emoji/sticker markup.
|
||||
func (s *Service) CreateAvatarVideoMarkupFromUpload(ctx context.Context, file domain.UploadedFileRef, videoStartTs float64, markup domain.PhotoSize) (domain.Photo, error) {
|
||||
if err := validateAvatarMarkupSize(markup); err != nil {
|
||||
return domain.Photo{}, err
|
||||
}
|
||||
return s.createAvatarVideoFromUpload(ctx, file, videoStartTs, []domain.PhotoSize{markup})
|
||||
}
|
||||
|
||||
func (s *Service) createAvatarVideoFromUpload(ctx context.Context, file domain.UploadedFileRef, videoStartTs float64, extraSizes []domain.PhotoSize) (domain.Photo, error) {
|
||||
body, err := s.assembleUploadBlob(ctx, file.OwnerUserID, file.FileID, file.Parts)
|
||||
if err != nil {
|
||||
return domain.Photo{}, err
|
||||
}
|
||||
if body.Size == 0 {
|
||||
return domain.Photo{}, domain.ErrPhotoInvalid
|
||||
}
|
||||
photoID := randomID()
|
||||
blob := domain.FileBlob{
|
||||
LocationKey: fmt.Sprintf("photo:%d:u", photoID),
|
||||
Backend: domain.MediaBackend(s.blobs.Name()),
|
||||
ObjectKey: body.ObjectKey,
|
||||
Size: body.Size,
|
||||
SHA256: body.SHA256,
|
||||
MimeType: "video/mp4",
|
||||
}
|
||||
if err := s.media.PutFileBlob(ctx, blob); err != nil {
|
||||
return domain.Photo{}, err
|
||||
}
|
||||
s.blobCache.put(blob.LocationKey, blob)
|
||||
stillBytes := s.avatarVideoStill(ctx, body, extraSizes)
|
||||
sizes, err := s.putPhotoStaticSizes(ctx, photoID, stillBytes, photoSizeSpecsForAvatar(stillBytes))
|
||||
if err != nil {
|
||||
return domain.Photo{}, err
|
||||
}
|
||||
sizes = append(sizes, domain.PhotoSize{
|
||||
Kind: domain.PhotoSizeKindVideo,
|
||||
Type: "u",
|
||||
W: 640,
|
||||
H: 640,
|
||||
Size: int(body.Size),
|
||||
VideoStartTs: videoStartTs,
|
||||
})
|
||||
sizes = append(sizes, extraSizes...)
|
||||
photo := domain.Photo{
|
||||
ID: photoID,
|
||||
AccessHash: randomID(),
|
||||
FileReference: randomFileReference(),
|
||||
Date: int(time.Now().Unix()),
|
||||
DCID: s.dc,
|
||||
Sizes: sizes,
|
||||
}
|
||||
if err := s.media.PutPhoto(ctx, photo); err != nil {
|
||||
return domain.Photo{}, err
|
||||
}
|
||||
if err := s.cleanupUploadParts(ctx, file.OwnerUserID, file.FileID); err != nil {
|
||||
s.log.Warn("cleanup assembled avatar video upload parts failed",
|
||||
zap.Int64("owner_user_id", file.OwnerUserID),
|
||||
zap.Int64("file_id", file.FileID),
|
||||
zap.Int64("photo_id", photoID),
|
||||
zap.Error(err))
|
||||
}
|
||||
return photo, nil
|
||||
}
|
||||
|
||||
// CreateAvatarMarkup stores an emoji/sticker animated profile markup as photo.video_sizes.
|
||||
func (s *Service) CreateAvatarMarkup(ctx context.Context, size domain.PhotoSize) (domain.Photo, error) {
|
||||
if err := validateAvatarMarkupSize(size); err != nil {
|
||||
return domain.Photo{}, err
|
||||
}
|
||||
photoID := randomID()
|
||||
stillBytes := s.generatedAvatarStill(ctx, size)
|
||||
sizes, err := s.putPhotoStaticSizes(ctx, photoID, stillBytes, photoSizeSpecsForAvatar(stillBytes))
|
||||
if err != nil {
|
||||
return domain.Photo{}, err
|
||||
}
|
||||
sizes = append(sizes, size)
|
||||
photo := domain.Photo{
|
||||
ID: photoID,
|
||||
AccessHash: randomID(),
|
||||
FileReference: randomFileReference(),
|
||||
Date: int(time.Now().Unix()),
|
||||
DCID: s.dc,
|
||||
Sizes: sizes,
|
||||
}
|
||||
if err := s.media.PutPhoto(ctx, photo); err != nil {
|
||||
return domain.Photo{}, err
|
||||
}
|
||||
return photo, nil
|
||||
}
|
||||
|
||||
func validateAvatarMarkupSize(size domain.PhotoSize) error {
|
||||
switch size.Kind {
|
||||
case domain.PhotoSizeKindVideoEmojiMarkup:
|
||||
if size.EmojiID == 0 || len(size.BackgroundColors) == 0 {
|
||||
return domain.ErrPhotoInvalid
|
||||
}
|
||||
case domain.PhotoSizeKindVideoStickerMarkup:
|
||||
if size.StickerID == 0 || len(size.BackgroundColors) == 0 {
|
||||
return domain.ErrPhotoInvalid
|
||||
}
|
||||
default:
|
||||
return domain.ErrPhotoInvalid
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// CreateDocumentFromUpload 把已上传文件组装成 Document(文件/视频/音频/gif/贴纸消息),落 blob + documents。
|
||||
func (s *Service) CreateDocumentFromUpload(ctx context.Context, file domain.UploadedFileRef, spec domain.DocumentSpec) (domain.Document, error) {
|
||||
data, err := s.assembleUpload(ctx, file.OwnerUserID, file.FileID, file.Parts)
|
||||
body, err := s.assembleUploadBlob(ctx, file.OwnerUserID, file.FileID, file.Parts)
|
||||
if err != nil {
|
||||
return domain.Document{}, err
|
||||
}
|
||||
if len(data) == 0 {
|
||||
if body.Size == 0 {
|
||||
return domain.Document{}, domain.ErrDocumentInvalid
|
||||
}
|
||||
objectKey, err := s.blobs.Put(ctx, data)
|
||||
if err != nil {
|
||||
return domain.Document{}, err
|
||||
}
|
||||
// faststart:MP4 视频若 moov 在末尾,搬到文件头以支持流式播放。普通 Telegram 客户端
|
||||
// 上传前会做这步;DrKLO 发 story 视频不转码导致 moov 在末尾,TDesktop 流式播放路径
|
||||
// 无法解复用(av_read_frame Invalid data)。不转码、保留原编码(含 HEVC)。
|
||||
body = s.maybeFaststartVideoBlob(ctx, spec.MimeType, body)
|
||||
docID := randomID()
|
||||
if err := s.media.PutFileBlob(ctx, domain.FileBlob{
|
||||
LocationKey: fmt.Sprintf("doc:%d", docID),
|
||||
Backend: domain.MediaBackend(s.blobs.Name()),
|
||||
ObjectKey: objectKey,
|
||||
Size: int64(len(data)),
|
||||
ObjectKey: body.ObjectKey,
|
||||
Size: body.Size,
|
||||
SHA256: body.SHA256,
|
||||
MimeType: spec.MimeType,
|
||||
}); err != nil {
|
||||
return domain.Document{}, err
|
||||
|
|
@ -108,28 +230,183 @@ func (s *Service) CreateDocumentFromUpload(ctx context.Context, file domain.Uplo
|
|||
FileReference: randomFileReference(),
|
||||
Date: int(time.Now().Unix()),
|
||||
MimeType: spec.MimeType,
|
||||
Size: int64(len(data)),
|
||||
Size: body.Size,
|
||||
DCID: s.dc,
|
||||
Attributes: spec.Attributes,
|
||||
}
|
||||
if spec.Thumb != nil {
|
||||
thumbData, err := s.assembleUpload(ctx, spec.Thumb.OwnerUserID, spec.Thumb.FileID, spec.Thumb.Parts)
|
||||
if err == nil && len(thumbData) > 0 {
|
||||
thumbKey, err := s.blobs.Put(ctx, thumbData)
|
||||
if err == nil {
|
||||
w, h := imageDimensions(thumbData, 0, 0)
|
||||
if err := s.media.PutFileBlob(ctx, domain.FileBlob{
|
||||
LocationKey: fmt.Sprintf("doc:%d:m", docID),
|
||||
Backend: domain.MediaBackend(s.blobs.Name()),
|
||||
ObjectKey: thumbKey,
|
||||
Size: int64(len(thumbData)),
|
||||
MimeType: "image/jpeg",
|
||||
}); err == nil {
|
||||
doc.Thumbs = []domain.PhotoSize{{Kind: domain.PhotoSizeKindDefault, Type: "m", W: w, H: h, Size: len(thumbData)}}
|
||||
}
|
||||
if thumb, err := s.putDocumentThumb(ctx, docID, thumbData); err == nil {
|
||||
doc.Thumbs = []domain.PhotoSize{thumb}
|
||||
}
|
||||
}
|
||||
}
|
||||
if len(doc.Thumbs) == 0 {
|
||||
if thumb, ok := s.generateVideoThumbFallbackFromBlob(ctx, docID, body.ObjectKey, body.Size, spec); ok {
|
||||
doc.Thumbs = []domain.PhotoSize{thumb}
|
||||
}
|
||||
}
|
||||
if err := s.media.PutDocument(ctx, doc); err != nil {
|
||||
return domain.Document{}, err
|
||||
}
|
||||
if err := s.cleanupUploadParts(ctx, file.OwnerUserID, file.FileID); err != nil {
|
||||
s.log.Warn("cleanup assembled document upload parts failed",
|
||||
zap.Int64("owner_user_id", file.OwnerUserID),
|
||||
zap.Int64("file_id", file.FileID),
|
||||
zap.Int64("document_id", docID),
|
||||
zap.Error(err))
|
||||
}
|
||||
return doc, nil
|
||||
}
|
||||
|
||||
// maxFaststartBytes 限制 faststart 一次性载入内存的视频大小;超过则跳过(流式 faststart
|
||||
// 复杂度高,超大视频是边角)。与缩略图回退路径同样全量读 blob,内存模式无新增。
|
||||
const maxFaststartBytes = 200 << 20
|
||||
|
||||
// faststartVideoMimes 是会用 moov/mdat 结构、值得尝试 faststart 的容器 mime。
|
||||
// 其它 video/*(webm 等)结构不同,faststartMP4 也会自检后 no-op,故不在此列以免白读 blob。
|
||||
var faststartVideoMimes = map[string]bool{
|
||||
"video/mp4": true,
|
||||
"video/quicktime": true,
|
||||
"video/x-m4v": true,
|
||||
}
|
||||
|
||||
// maybeFaststartVideoBlob 对 MP4/MOV 视频上传尝试 faststart。性能考量:
|
||||
// 1. 先只读顶层 box 头(廉价探测)判断是否需要——绝大多数客户端上传的视频本就已 faststart,
|
||||
// 此时只发生几次 16 字节读,不读整段媒体。
|
||||
// 2. 仅 moov 在末尾时才重写;且优先走流式(仅 ftyp+moov 进内存,mdat 大块分块流式拼接),
|
||||
// 不把整段视频 2× 驻留内存。moov 非末尾的罕见排布回退到全量重排。
|
||||
// 任何不适用/失败都返回原 body,绝不让上传失败或损坏数据。
|
||||
func (s *Service) maybeFaststartVideoBlob(ctx context.Context, mimeType string, body assembledUploadBlob) assembledUploadBlob {
|
||||
if !faststartVideoMimes[strings.ToLower(strings.TrimSpace(mimeType))] {
|
||||
return body
|
||||
}
|
||||
if body.Size <= 0 || body.Size > maxFaststartBytes {
|
||||
return body
|
||||
}
|
||||
readAt := func(off, n int64) ([]byte, error) {
|
||||
data, _, err := s.blobs.GetRange(ctx, body.ObjectKey, off, n)
|
||||
return data, err
|
||||
}
|
||||
layout, ok := inspectMP4Layout(body.Size, readAt)
|
||||
if !ok || !layout.needsFaststart {
|
||||
return body // 非 MP4 / 已 faststart —— 未读整段媒体
|
||||
}
|
||||
|
||||
var reader io.Reader
|
||||
if layout.moovIsLast {
|
||||
// 流式重写:只把 ftyp + moov 读进内存并 patch 偏移,mdat 区段分块流式。
|
||||
ftyp, e1 := readAt(layout.ftypStart, layout.ftypEnd-layout.ftypStart)
|
||||
moov, e2 := readAt(layout.moovStart, layout.moovEnd-layout.moovStart)
|
||||
moovSize := layout.moovEnd - layout.moovStart
|
||||
if e1 != nil || e2 != nil ||
|
||||
int64(len(ftyp)) != layout.ftypEnd-layout.ftypStart ||
|
||||
int64(len(moov)) != moovSize ||
|
||||
!patchChunkOffsets(moov, moovSize) {
|
||||
return body
|
||||
}
|
||||
mid := &blobRangeReader{ctx: ctx, blobs: s.blobs, key: body.ObjectKey, pos: layout.ftypEnd, end: layout.moovStart}
|
||||
reader = io.MultiReader(bytes.NewReader(ftyp), bytes.NewReader(moov), mid)
|
||||
} else {
|
||||
// 罕见:moov 非末尾。回退到全量读 + 重排(已测函数)。
|
||||
data, total, err := s.blobs.GetRange(ctx, body.ObjectKey, 0, body.Size)
|
||||
if err != nil || total != body.Size || int64(len(data)) != body.Size {
|
||||
return body
|
||||
}
|
||||
out, changed := faststartMP4(data)
|
||||
if !changed {
|
||||
return body
|
||||
}
|
||||
reader = bytes.NewReader(out)
|
||||
}
|
||||
|
||||
key, size, sum, err := s.blobs.PutReader(ctx, reader)
|
||||
if err != nil {
|
||||
s.log.Warn("faststart re-store failed; keeping original blob",
|
||||
zap.String("mime", mimeType), zap.Int64("size", body.Size), zap.Error(err))
|
||||
return body
|
||||
}
|
||||
if size != body.Size {
|
||||
// faststart 守恒大小;不等说明流式拼接出错,丢弃新 blob 用原 blob 兜底。
|
||||
s.log.Warn("faststart size mismatch; keeping original blob",
|
||||
zap.Int64("orig", body.Size), zap.Int64("got", size))
|
||||
return body
|
||||
}
|
||||
s.log.Info("faststart applied to uploaded video",
|
||||
zap.String("mime", mimeType), zap.Int64("size", size))
|
||||
return assembledUploadBlob{ObjectKey: key, Size: size, SHA256: sum}
|
||||
}
|
||||
|
||||
// blobRangeReader 把 blob 的 [pos, end) 区段按 io.Reader 调用方给的缓冲大小分块流式读出,
|
||||
// 用于 faststart 流式拼接 mdat,避免整段媒体驻留内存。
|
||||
type blobRangeReader struct {
|
||||
ctx context.Context
|
||||
blobs BlobBackend
|
||||
key string
|
||||
pos int64
|
||||
end int64
|
||||
}
|
||||
|
||||
func (r *blobRangeReader) Read(p []byte) (int, error) {
|
||||
if r.pos >= r.end {
|
||||
return 0, io.EOF
|
||||
}
|
||||
want := r.end - r.pos
|
||||
if want > int64(len(p)) {
|
||||
want = int64(len(p))
|
||||
}
|
||||
data, _, err := r.blobs.GetRange(r.ctx, r.key, r.pos, want)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
if len(data) == 0 {
|
||||
return 0, io.ErrUnexpectedEOF // 区段内不应读到空,避免静默截断
|
||||
}
|
||||
n := copy(p, data)
|
||||
r.pos += int64(n)
|
||||
return n, nil
|
||||
}
|
||||
|
||||
// CreateDocumentFromBytes stores already-fetched bytes as a message Document.
|
||||
func (s *Service) CreateDocumentFromBytes(ctx context.Context, data []byte, spec domain.DocumentSpec) (domain.Document, error) {
|
||||
if len(data) == 0 {
|
||||
return domain.Document{}, domain.ErrDocumentInvalid
|
||||
}
|
||||
if strings.TrimSpace(spec.MimeType) == "" {
|
||||
spec.MimeType = "application/octet-stream"
|
||||
}
|
||||
objectKey, size, sum, err := s.blobs.PutReader(ctx, bytes.NewReader(data))
|
||||
if err != nil {
|
||||
return domain.Document{}, err
|
||||
}
|
||||
if size == 0 {
|
||||
return domain.Document{}, domain.ErrDocumentInvalid
|
||||
}
|
||||
docID := randomID()
|
||||
if err := s.media.PutFileBlob(ctx, domain.FileBlob{
|
||||
LocationKey: fmt.Sprintf("doc:%d", docID),
|
||||
Backend: domain.MediaBackend(s.blobs.Name()),
|
||||
ObjectKey: objectKey,
|
||||
Size: size,
|
||||
SHA256: sum,
|
||||
MimeType: spec.MimeType,
|
||||
}); err != nil {
|
||||
return domain.Document{}, err
|
||||
}
|
||||
doc := domain.Document{
|
||||
ID: docID,
|
||||
AccessHash: randomID(),
|
||||
FileReference: randomFileReference(),
|
||||
Date: int(time.Now().Unix()),
|
||||
MimeType: spec.MimeType,
|
||||
Size: size,
|
||||
DCID: s.dc,
|
||||
Attributes: spec.Attributes,
|
||||
}
|
||||
if thumb, ok := s.generateVideoThumbFallbackFromBlob(ctx, docID, objectKey, size, spec); ok {
|
||||
doc.Thumbs = []domain.PhotoSize{thumb}
|
||||
}
|
||||
if err := s.media.PutDocument(ctx, doc); err != nil {
|
||||
return domain.Document{}, err
|
||||
}
|
||||
|
|
@ -177,19 +454,7 @@ func (s *Service) GetProfilePhotos(ctx context.Context, ownerType domain.PeerTyp
|
|||
|
||||
// GetProfilePhotosKind returns profile/fallback photo history.
|
||||
func (s *Service) GetProfilePhotosKind(ctx context.Context, ownerType domain.PeerType, ownerID int64, kind domain.ProfilePhotoKind, offset, limit int, maxID int64) ([]domain.Photo, int, error) {
|
||||
ids, total, err := s.media.ListProfilePhotosKind(ctx, ownerType, ownerID, kind, offset, limit, maxID)
|
||||
if err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
photos := make([]domain.Photo, 0, len(ids))
|
||||
for _, id := range ids {
|
||||
if p, ok, err := s.media.GetPhoto(ctx, id); err != nil {
|
||||
return nil, 0, err
|
||||
} else if ok {
|
||||
photos = append(photos, p)
|
||||
}
|
||||
}
|
||||
return photos, total, nil
|
||||
return s.media.ListProfilePhotoDetailsKind(ctx, ownerType, ownerID, kind, offset, limit, maxID)
|
||||
}
|
||||
|
||||
// DeleteProfilePhotos 停用指定头像,返回成功停用数量。
|
||||
|
|
@ -208,24 +473,11 @@ func (s *Service) DeleteProfilePhotosKind(ctx context.Context, ownerType domain.
|
|||
|
||||
// createPhoto 把字节落 blob(每个尺寸一个 location_key,指向同一内容)并写 photos 表。
|
||||
func (s *Service) createPhoto(ctx context.Context, data []byte, specs []photoSizeSpec) (domain.Photo, error) {
|
||||
objectKey, err := s.blobs.Put(ctx, data)
|
||||
photoID := randomID()
|
||||
sizes, err := s.putPhotoStaticSizes(ctx, photoID, data, specs)
|
||||
if err != nil {
|
||||
return domain.Photo{}, err
|
||||
}
|
||||
photoID := randomID()
|
||||
sizes := make([]domain.PhotoSize, 0, len(specs))
|
||||
for _, spec := range specs {
|
||||
if err := s.media.PutFileBlob(ctx, domain.FileBlob{
|
||||
LocationKey: fmt.Sprintf("photo:%d:%s", photoID, spec.Type),
|
||||
Backend: domain.MediaBackend(s.blobs.Name()),
|
||||
ObjectKey: objectKey,
|
||||
Size: int64(len(data)),
|
||||
MimeType: "image/jpeg",
|
||||
}); err != nil {
|
||||
return domain.Photo{}, err
|
||||
}
|
||||
sizes = append(sizes, domain.PhotoSize{Kind: domain.PhotoSizeKindDefault, Type: spec.Type, W: spec.W, H: spec.H, Size: len(data)})
|
||||
}
|
||||
photo := domain.Photo{
|
||||
ID: photoID,
|
||||
AccessHash: randomID(),
|
||||
|
|
@ -240,6 +492,116 @@ func (s *Service) createPhoto(ctx context.Context, data []byte, specs []photoSiz
|
|||
return photo, nil
|
||||
}
|
||||
|
||||
func (s *Service) putPhotoStaticSizes(ctx context.Context, photoID int64, data []byte, specs []photoSizeSpec) ([]domain.PhotoSize, error) {
|
||||
objectKey, err := s.blobs.Put(ctx, data)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
mimeType := imageMimeType(data)
|
||||
sizes := make([]domain.PhotoSize, 0, len(specs))
|
||||
for _, spec := range specs {
|
||||
blob := domain.FileBlob{
|
||||
LocationKey: fmt.Sprintf("photo:%d:%s", photoID, spec.Type),
|
||||
Backend: domain.MediaBackend(s.blobs.Name()),
|
||||
ObjectKey: objectKey,
|
||||
Size: int64(len(data)),
|
||||
MimeType: mimeType,
|
||||
}
|
||||
if err := s.media.PutFileBlob(ctx, blob); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
s.blobCache.put(blob.LocationKey, blob)
|
||||
sizes = append(sizes, domain.PhotoSize{Kind: domain.PhotoSizeKindDefault, Type: spec.Type, W: spec.W, H: spec.H, Size: len(data)})
|
||||
}
|
||||
s.prewarmSmallBlob(objectKey, data)
|
||||
return sizes, nil
|
||||
}
|
||||
|
||||
func (s *Service) putDocumentThumb(ctx context.Context, docID int64, thumbData []byte) (domain.PhotoSize, error) {
|
||||
if len(thumbData) == 0 {
|
||||
return domain.PhotoSize{}, fmt.Errorf("empty document thumbnail")
|
||||
}
|
||||
thumbKey, err := s.blobs.Put(ctx, thumbData)
|
||||
if err != nil {
|
||||
return domain.PhotoSize{}, err
|
||||
}
|
||||
w, h := imageDimensions(thumbData, 0, 0)
|
||||
blob := domain.FileBlob{
|
||||
LocationKey: fmt.Sprintf("doc:%d:m", docID),
|
||||
Backend: domain.MediaBackend(s.blobs.Name()),
|
||||
ObjectKey: thumbKey,
|
||||
Size: int64(len(thumbData)),
|
||||
MimeType: "image/jpeg",
|
||||
}
|
||||
if err := s.media.PutFileBlob(ctx, blob); err != nil {
|
||||
return domain.PhotoSize{}, err
|
||||
}
|
||||
s.blobCache.put(blob.LocationKey, blob)
|
||||
s.prewarmSmallBlob(thumbKey, thumbData)
|
||||
return domain.PhotoSize{Kind: domain.PhotoSizeKindDefault, Type: "m", W: w, H: h, Size: len(thumbData)}, nil
|
||||
}
|
||||
|
||||
func (s *Service) generateVideoThumbFallback(ctx context.Context, docID int64, data []byte, spec domain.DocumentSpec) (domain.PhotoSize, bool) {
|
||||
if s.thumbs == nil || !documentSpecIsVideo(spec) || len(data) > videoThumbnailMaxInputBytes {
|
||||
return domain.PhotoSize{}, false
|
||||
}
|
||||
thumbData, err := s.thumbs.Extract(ctx, data, spec.MimeType)
|
||||
if err != nil {
|
||||
s.log.Warn("server-side video thumbnail fallback failed",
|
||||
zap.Int64("document_id", docID),
|
||||
zap.String("mime_type", spec.MimeType),
|
||||
zap.Int64("bytes", int64(len(data))),
|
||||
zap.Error(err))
|
||||
return domain.PhotoSize{}, false
|
||||
}
|
||||
thumb, err := s.putDocumentThumb(ctx, docID, thumbData)
|
||||
if err != nil {
|
||||
s.log.Warn("store server-side video thumbnail failed",
|
||||
zap.Int64("document_id", docID),
|
||||
zap.String("mime_type", spec.MimeType),
|
||||
zap.Int64("thumb_bytes", int64(len(thumbData))),
|
||||
zap.Error(err))
|
||||
return domain.PhotoSize{}, false
|
||||
}
|
||||
return thumb, true
|
||||
}
|
||||
|
||||
func (s *Service) generateVideoThumbFallbackFromBlob(ctx context.Context, docID int64, objectKey string, size int64, spec domain.DocumentSpec) (domain.PhotoSize, bool) {
|
||||
if s.thumbs == nil || !documentSpecIsVideo(spec) || size > videoThumbnailMaxInputBytes {
|
||||
return domain.PhotoSize{}, false
|
||||
}
|
||||
data, total, err := s.blobs.GetRange(ctx, objectKey, 0, size)
|
||||
if err != nil {
|
||||
s.log.Warn("read video blob for thumbnail fallback failed",
|
||||
zap.Int64("document_id", docID),
|
||||
zap.String("mime_type", spec.MimeType),
|
||||
zap.Int64("bytes", size),
|
||||
zap.Error(err))
|
||||
return domain.PhotoSize{}, false
|
||||
}
|
||||
if int64(len(data)) != total || total != size {
|
||||
s.log.Warn("video blob size mismatch for thumbnail fallback",
|
||||
zap.Int64("document_id", docID),
|
||||
zap.Int64("expected_size", size),
|
||||
zap.Int64("total_size", total),
|
||||
zap.Int("read_bytes", len(data)))
|
||||
return domain.PhotoSize{}, false
|
||||
}
|
||||
return s.generateVideoThumbFallback(ctx, docID, data, spec)
|
||||
}
|
||||
|
||||
func documentSpecIsVideo(spec domain.DocumentSpec) bool {
|
||||
if strings.HasPrefix(strings.ToLower(strings.TrimSpace(spec.MimeType)), "video/") {
|
||||
return true
|
||||
}
|
||||
for _, attr := range spec.Attributes {
|
||||
if attr.Kind == domain.DocAttrVideo {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
type photoSizeSpec struct {
|
||||
Type string
|
||||
W int
|
||||
|
|
@ -286,6 +648,281 @@ func scaleDown(w, h, max int) (int, int) {
|
|||
return max * w / h, max
|
||||
}
|
||||
|
||||
const (
|
||||
avatarStillSize = 640
|
||||
avatarMarkupScale = 0.70
|
||||
avatarMarkupMaxSourceBytes = 2 << 20 // emoji/sticker thumb 小对象保护线。
|
||||
)
|
||||
|
||||
// avatarVideoStill 生成动画头像的静态尺寸字节:优先抽取上传视频首帧——动画头像
|
||||
// (emoji/sticker 构造器或自选视频)的首帧就是用户在客户端看到的真实画面(彩色
|
||||
// emoji、圆角、布局都一致);抽帧不可用时回退到按 markup 服务端合成。
|
||||
func (s *Service) avatarVideoStill(ctx context.Context, body assembledUploadBlob, extraSizes []domain.PhotoSize) []byte {
|
||||
if s.thumbs != nil && body.Size > 0 && body.Size <= videoThumbnailMaxInputBytes {
|
||||
data, total, err := s.blobs.GetRange(ctx, body.ObjectKey, 0, body.Size)
|
||||
if err == nil && int64(len(data)) == total && total == body.Size {
|
||||
if thumb, err := s.thumbs.Extract(ctx, data, "video/mp4"); err == nil && len(thumb) > 0 {
|
||||
return thumb
|
||||
} else if err != nil {
|
||||
s.log.Debug("extract avatar video first frame failed, falling back to composed still",
|
||||
zap.String("object_key", body.ObjectKey),
|
||||
zap.Int64("bytes", body.Size),
|
||||
zap.Error(err))
|
||||
}
|
||||
} else if err != nil {
|
||||
s.log.Warn("read avatar video blob for still failed",
|
||||
zap.String("object_key", body.ObjectKey),
|
||||
zap.Int64("bytes", body.Size),
|
||||
zap.Error(err))
|
||||
}
|
||||
}
|
||||
return s.generatedAvatarStill(ctx, avatarStillMarkup(extraSizes))
|
||||
}
|
||||
|
||||
func (s *Service) generatedAvatarStill(ctx context.Context, markup domain.PhotoSize) []byte {
|
||||
img := generatedAvatarBackground(markup.BackgroundColors)
|
||||
if overlay, tintWhite, ok := s.avatarMarkupOverlay(ctx, markup); ok {
|
||||
drawAvatarMarkup(img, overlay, tintWhite)
|
||||
}
|
||||
return encodeAvatarPNG(img)
|
||||
}
|
||||
|
||||
func generatedAvatarBackground(colors []int) *image.RGBA {
|
||||
if len(colors) == 0 {
|
||||
colors = []int{0x5b8def, 0x53c6a4}
|
||||
}
|
||||
first := rgbColor(colors[0])
|
||||
last := first
|
||||
if len(colors) > 1 {
|
||||
last = rgbColor(colors[len(colors)-1])
|
||||
}
|
||||
img := image.NewRGBA(image.Rect(0, 0, avatarStillSize, avatarStillSize))
|
||||
for y := 0; y < avatarStillSize; y++ {
|
||||
t := float64(y) / float64(avatarStillSize-1)
|
||||
row := lerpColor(first, last, t)
|
||||
for x := 0; x < avatarStillSize; x++ {
|
||||
img.SetRGBA(x, y, row)
|
||||
}
|
||||
}
|
||||
return img
|
||||
}
|
||||
|
||||
func encodeAvatarPNG(img image.Image) []byte {
|
||||
var buf bytes.Buffer
|
||||
_ = png.Encode(&buf, img)
|
||||
return buf.Bytes()
|
||||
}
|
||||
|
||||
// avatarMarkupOverlay 加载 markup 引用文档的静态缩略图作为合成贴图。第二个返回值
|
||||
// 表示是否染白:仅 text_color 的 custom emoji(单色、由客户端按文字色适配渲染,
|
||||
// 头像背景上的约定呈现是白色剪影)需要染白,普通彩色 emoji / sticker 保留原色。
|
||||
func (s *Service) avatarMarkupOverlay(ctx context.Context, markup domain.PhotoSize) (image.Image, bool, bool) {
|
||||
if s == nil || s.media == nil {
|
||||
return nil, false, false
|
||||
}
|
||||
docID := int64(0)
|
||||
switch markup.Kind {
|
||||
case domain.PhotoSizeKindVideoEmojiMarkup:
|
||||
docID = markup.EmojiID
|
||||
case domain.PhotoSizeKindVideoStickerMarkup:
|
||||
docID = markup.StickerID
|
||||
default:
|
||||
return nil, false, false
|
||||
}
|
||||
if docID == 0 {
|
||||
return nil, false, false
|
||||
}
|
||||
doc, found, err := s.media.GetDocument(ctx, docID)
|
||||
if err != nil {
|
||||
s.log.Warn("load avatar markup document failed", zap.Int64("document_id", docID), zap.Error(err))
|
||||
return nil, false, false
|
||||
}
|
||||
if !found {
|
||||
return nil, false, false
|
||||
}
|
||||
data, ok := s.avatarMarkupBytes(ctx, doc)
|
||||
if !ok {
|
||||
return nil, false, false
|
||||
}
|
||||
img, _, err := image.Decode(bytes.NewReader(data))
|
||||
if err != nil {
|
||||
s.log.Debug("decode avatar markup thumbnail failed",
|
||||
zap.Int64("document_id", doc.ID),
|
||||
zap.String("mime_type", doc.MimeType),
|
||||
zap.Error(err))
|
||||
return nil, false, false
|
||||
}
|
||||
return img, documentIsTextColorEmoji(doc), true
|
||||
}
|
||||
|
||||
// documentIsTextColorEmoji 判断文档是否为声明 text_color 的 custom emoji。
|
||||
func documentIsTextColorEmoji(doc domain.Document) bool {
|
||||
for _, attr := range doc.Attributes {
|
||||
if attr.Kind == domain.DocAttrCustomEmoji {
|
||||
return attr.TextColor
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func (s *Service) avatarMarkupBytes(ctx context.Context, doc domain.Document) ([]byte, bool) {
|
||||
var best []byte
|
||||
bestScore := -1
|
||||
for _, thumb := range doc.Thumbs {
|
||||
score := avatarThumbScore(thumb)
|
||||
if score <= bestScore {
|
||||
continue
|
||||
}
|
||||
switch thumb.Kind {
|
||||
case domain.PhotoSizeKindCached:
|
||||
if len(thumb.Bytes) == 0 || len(thumb.Bytes) > avatarMarkupMaxSourceBytes {
|
||||
continue
|
||||
}
|
||||
best = append([]byte(nil), thumb.Bytes...)
|
||||
bestScore = score
|
||||
case domain.PhotoSizeKindDefault, domain.PhotoSizeKindProgressive:
|
||||
if thumb.Type == "" || thumb.Size <= 0 || thumb.Size > avatarMarkupMaxSourceBytes {
|
||||
continue
|
||||
}
|
||||
data, ok := s.readSmallBlob(ctx, fmt.Sprintf("doc:%d:%s", doc.ID, thumb.Type), int64(thumb.Size))
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
best = data
|
||||
bestScore = score
|
||||
}
|
||||
}
|
||||
if len(best) > 0 {
|
||||
return best, true
|
||||
}
|
||||
if strings.HasPrefix(strings.ToLower(strings.TrimSpace(doc.MimeType)), "image/") &&
|
||||
doc.Size > 0 && doc.Size <= avatarMarkupMaxSourceBytes {
|
||||
return s.readSmallBlob(ctx, fmt.Sprintf("doc:%d", doc.ID), doc.Size)
|
||||
}
|
||||
return nil, false
|
||||
}
|
||||
|
||||
func (s *Service) readSmallBlob(ctx context.Context, locationKey string, expectedSize int64) ([]byte, bool) {
|
||||
if expectedSize <= 0 || expectedSize > avatarMarkupMaxSourceBytes {
|
||||
return nil, false
|
||||
}
|
||||
blob, found, err := s.media.GetFileBlob(ctx, locationKey)
|
||||
if err != nil {
|
||||
s.log.Warn("load avatar markup blob metadata failed",
|
||||
zap.String("location_key", locationKey),
|
||||
zap.Error(err))
|
||||
return nil, false
|
||||
}
|
||||
if !found || blob.Size <= 0 || blob.Size > avatarMarkupMaxSourceBytes {
|
||||
return nil, false
|
||||
}
|
||||
data, total, err := s.blobs.GetRange(ctx, blob.ObjectKey, 0, blob.Size)
|
||||
if err != nil {
|
||||
s.log.Warn("read avatar markup blob failed",
|
||||
zap.String("location_key", locationKey),
|
||||
zap.String("object_key", blob.ObjectKey),
|
||||
zap.Error(err))
|
||||
return nil, false
|
||||
}
|
||||
if int64(len(data)) != total || total != blob.Size {
|
||||
return nil, false
|
||||
}
|
||||
return data, true
|
||||
}
|
||||
|
||||
func avatarThumbScore(size domain.PhotoSize) int {
|
||||
if size.W > 0 && size.H > 0 {
|
||||
return size.W * size.H
|
||||
}
|
||||
if size.Size > 0 {
|
||||
return size.Size
|
||||
}
|
||||
return len(size.Bytes)
|
||||
}
|
||||
|
||||
// drawAvatarMarkup 把贴图缩放后居中画到背景上。默认保留贴图原色(彩色 emoji /
|
||||
// sticker 头像),仅 tintWhite 时染成白色剪影。
|
||||
func drawAvatarMarkup(dst *image.RGBA, src image.Image, tintWhite bool) {
|
||||
srcBounds := src.Bounds()
|
||||
sw, sh := srcBounds.Dx(), srcBounds.Dy()
|
||||
if sw <= 0 || sh <= 0 {
|
||||
return
|
||||
}
|
||||
max := int(math.Round(float64(dst.Bounds().Dx()) * avatarMarkupScale))
|
||||
if max <= 0 {
|
||||
return
|
||||
}
|
||||
scale := math.Min(float64(max)/float64(sw), float64(max)/float64(sh))
|
||||
w := maxInt(1, int(math.Round(float64(sw)*scale)))
|
||||
h := maxInt(1, int(math.Round(float64(sh)*scale)))
|
||||
rect := image.Rect(
|
||||
(dst.Bounds().Dx()-w)/2,
|
||||
(dst.Bounds().Dy()-h)/2,
|
||||
(dst.Bounds().Dx()+w)/2,
|
||||
(dst.Bounds().Dy()+h)/2,
|
||||
)
|
||||
scaled := image.NewRGBA(image.Rect(0, 0, rect.Dx(), rect.Dy()))
|
||||
xdraw.CatmullRom.Scale(scaled, scaled.Bounds(), src, srcBounds, xdraw.Src, nil)
|
||||
if tintWhite {
|
||||
tintWhitePremultiplied(scaled)
|
||||
}
|
||||
stddraw.Draw(dst, rect, scaled, image.Point{}, stddraw.Over)
|
||||
}
|
||||
|
||||
// tintWhitePremultiplied 把贴图就地染成白色剪影(保留 alpha 形状)。image.RGBA 是
|
||||
// alpha 预乘存储,分量必须满足 R,G,B ≤ A:若写入 R=G=B=255 而 A<255 的非法值,
|
||||
// draw.Over 合成会算术溢出回绕,凡 alpha 不恰为 255 的像素整片输出近黑色。
|
||||
func tintWhitePremultiplied(img *image.RGBA) {
|
||||
for y := img.Bounds().Min.Y; y < img.Bounds().Max.Y; y++ {
|
||||
for x := img.Bounds().Min.X; x < img.Bounds().Max.X; x++ {
|
||||
_, _, _, a := img.At(x, y).RGBA()
|
||||
v := uint8(a >> 8)
|
||||
img.SetRGBA(x, y, color.RGBA{R: v, G: v, B: v, A: v})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func maxInt(a, b int) int {
|
||||
if a > b {
|
||||
return a
|
||||
}
|
||||
return b
|
||||
}
|
||||
|
||||
func rgbColor(v int) color.RGBA {
|
||||
u := uint32(v)
|
||||
return color.RGBA{R: uint8(u >> 16), G: uint8(u >> 8), B: uint8(u), A: 255}
|
||||
}
|
||||
|
||||
func lerpColor(a, b color.RGBA, t float64) color.RGBA {
|
||||
lerp := func(x, y uint8) uint8 {
|
||||
return uint8(float64(x)*(1-t) + float64(y)*t)
|
||||
}
|
||||
return color.RGBA{R: lerp(a.R, b.R), G: lerp(a.G, b.G), B: lerp(a.B, b.B), A: 255}
|
||||
}
|
||||
|
||||
func imageMimeType(data []byte) string {
|
||||
if len(data) >= 8 && bytes.Equal(data[:8], []byte{0x89, 'P', 'N', 'G', '\r', '\n', 0x1a, '\n'}) {
|
||||
return "image/png"
|
||||
}
|
||||
if len(data) >= 3 && data[0] == 0xff && data[1] == 0xd8 && data[2] == 0xff {
|
||||
return "image/jpeg"
|
||||
}
|
||||
return "application/octet-stream"
|
||||
}
|
||||
|
||||
func avatarStillMarkup(sizes []domain.PhotoSize) domain.PhotoSize {
|
||||
for _, size := range sizes {
|
||||
if size.Kind == domain.PhotoSizeKindVideoEmojiMarkup || size.Kind == domain.PhotoSizeKindVideoStickerMarkup {
|
||||
return size
|
||||
}
|
||||
if len(size.BackgroundColors) > 0 {
|
||||
return domain.PhotoSize{BackgroundColors: append([]int(nil), size.BackgroundColors...)}
|
||||
}
|
||||
}
|
||||
return domain.PhotoSize{}
|
||||
}
|
||||
|
||||
func randomID() int64 {
|
||||
var b [8]byte
|
||||
_, _ = rand.Read(b[:])
|
||||
|
|
|
|||
466
internal/app/files/photos_test.go
Normal file
466
internal/app/files/photos_test.go
Normal file
|
|
@ -0,0 +1,466 @@
|
|||
package files
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"image"
|
||||
"image/color"
|
||||
"image/jpeg"
|
||||
"image/png"
|
||||
"testing"
|
||||
|
||||
"telesrv/internal/domain"
|
||||
)
|
||||
|
||||
func TestCreateDocumentFromUploadGeneratesVideoThumbWhenMissing(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
media := newFakeMediaStore()
|
||||
blobs, err := NewLocalFS(t.TempDir())
|
||||
if err != nil {
|
||||
t.Fatalf("NewLocalFS: %v", err)
|
||||
}
|
||||
thumbBytes := testJPEG(t, 4, 2)
|
||||
thumbnailer := &fakeVideoThumbnailer{thumb: thumbBytes}
|
||||
svc := NewService(media, blobs, 2, WithVideoThumbnailer(thumbnailer))
|
||||
|
||||
if _, err := svc.SaveFilePart(ctx, 10, 100, 0, []byte("fake-video-bytes")); err != nil {
|
||||
t.Fatalf("SaveFilePart: %v", err)
|
||||
}
|
||||
doc, err := svc.CreateDocumentFromUpload(ctx,
|
||||
domain.UploadedFileRef{OwnerUserID: 10, FileID: 100, Parts: 1, Name: "video.mp4"},
|
||||
domain.DocumentSpec{
|
||||
MimeType: "video/mp4",
|
||||
Attributes: []domain.DocumentAttribute{{Kind: domain.DocAttrVideo, W: 640, H: 360, Duration: 1}},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CreateDocumentFromUpload: %v", err)
|
||||
}
|
||||
if thumbnailer.calls != 1 {
|
||||
t.Fatalf("thumbnailer calls = %d, want 1", thumbnailer.calls)
|
||||
}
|
||||
if len(doc.Thumbs) != 1 {
|
||||
t.Fatalf("thumbs = %+v, want one generated thumbnail", doc.Thumbs)
|
||||
}
|
||||
if got := doc.Thumbs[0]; got.Type != "m" || got.W != 4 || got.H != 2 || got.Size != len(thumbBytes) {
|
||||
t.Fatalf("thumb = %+v, want m 4x2 size=%d", got, len(thumbBytes))
|
||||
}
|
||||
blob, ok, err := media.GetFileBlob(ctx, fmt.Sprintf("doc:%d:m", doc.ID))
|
||||
if err != nil || !ok {
|
||||
t.Fatalf("generated thumb blob ok=%v err=%v", ok, err)
|
||||
}
|
||||
gotBytes, err := blobs.Get(ctx, blob.ObjectKey)
|
||||
if err != nil {
|
||||
t.Fatalf("read generated thumb blob: %v", err)
|
||||
}
|
||||
if !bytes.Equal(gotBytes, thumbBytes) {
|
||||
t.Fatalf("generated thumb bytes mismatch")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCreateDocumentFromUploadKeepsClientThumb(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
media := newFakeMediaStore()
|
||||
blobs, err := NewLocalFS(t.TempDir())
|
||||
if err != nil {
|
||||
t.Fatalf("NewLocalFS: %v", err)
|
||||
}
|
||||
thumbnailer := &fakeVideoThumbnailer{err: errors.New("should not be called")}
|
||||
svc := NewService(media, blobs, 2, WithVideoThumbnailer(thumbnailer))
|
||||
clientThumb := testJPEG(t, 3, 5)
|
||||
|
||||
if _, err := svc.SaveFilePart(ctx, 10, 200, 0, []byte("fake-video-bytes")); err != nil {
|
||||
t.Fatalf("SaveFilePart video: %v", err)
|
||||
}
|
||||
if _, err := svc.SaveFilePart(ctx, 10, 201, 0, clientThumb); err != nil {
|
||||
t.Fatalf("SaveFilePart thumb: %v", err)
|
||||
}
|
||||
doc, err := svc.CreateDocumentFromUpload(ctx,
|
||||
domain.UploadedFileRef{OwnerUserID: 10, FileID: 200, Parts: 1, Name: "video.mp4"},
|
||||
domain.DocumentSpec{
|
||||
MimeType: "video/mp4",
|
||||
Attributes: []domain.DocumentAttribute{{Kind: domain.DocAttrVideo, W: 640, H: 360, Duration: 1}},
|
||||
Thumb: &domain.UploadedFileRef{OwnerUserID: 10, FileID: 201, Parts: 1, Name: "thumb.jpg"},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CreateDocumentFromUpload: %v", err)
|
||||
}
|
||||
if thumbnailer.calls != 0 {
|
||||
t.Fatalf("thumbnailer calls = %d, want 0 when client thumb is available", thumbnailer.calls)
|
||||
}
|
||||
if len(doc.Thumbs) != 1 {
|
||||
t.Fatalf("thumbs = %+v, want client thumbnail", doc.Thumbs)
|
||||
}
|
||||
if got := doc.Thumbs[0]; got.W != 3 || got.H != 5 || got.Size != len(clientThumb) {
|
||||
t.Fatalf("thumb = %+v, want client thumb dimensions", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCreateDocumentFromUploadWithoutThumbnailerDoesNotBlockVideo(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
media := newFakeMediaStore()
|
||||
blobs, err := NewLocalFS(t.TempDir())
|
||||
if err != nil {
|
||||
t.Fatalf("NewLocalFS: %v", err)
|
||||
}
|
||||
svc := NewService(media, blobs, 2, WithVideoThumbnailer(nil))
|
||||
|
||||
if _, err := svc.SaveFilePart(ctx, 10, 300, 0, []byte("fake-video-bytes")); err != nil {
|
||||
t.Fatalf("SaveFilePart: %v", err)
|
||||
}
|
||||
doc, err := svc.CreateDocumentFromUpload(ctx,
|
||||
domain.UploadedFileRef{OwnerUserID: 10, FileID: 300, Parts: 1, Name: "video.mp4"},
|
||||
domain.DocumentSpec{
|
||||
MimeType: "video/mp4",
|
||||
Attributes: []domain.DocumentAttribute{{Kind: domain.DocAttrVideo, W: 640, H: 360, Duration: 1}},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CreateDocumentFromUpload without thumbnailer: %v", err)
|
||||
}
|
||||
if len(doc.Thumbs) != 0 {
|
||||
t.Fatalf("thumbs = %+v, want no fallback thumbnail when thumbnailer is disabled", doc.Thumbs)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCreatePhotoFromBytesStoresDownloadableMessageSizes(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
media := newFakeMediaStore()
|
||||
blobs, err := NewLocalFS(t.TempDir())
|
||||
if err != nil {
|
||||
t.Fatalf("NewLocalFS: %v", err)
|
||||
}
|
||||
svc := NewService(media, blobs, 2)
|
||||
data := testJPEG(t, 16, 9)
|
||||
|
||||
photo, err := svc.CreatePhotoFromBytes(ctx, data)
|
||||
if err != nil {
|
||||
t.Fatalf("CreatePhotoFromBytes: %v", err)
|
||||
}
|
||||
if photo.ID == 0 || photo.AccessHash == 0 || photo.DCID != 2 || len(photo.Sizes) != 2 {
|
||||
t.Fatalf("photo = %+v, want stored photo with message sizes", photo)
|
||||
}
|
||||
blob, ok, err := media.GetFileBlob(ctx, fmt.Sprintf("photo:%d:x", photo.ID))
|
||||
if err != nil || !ok {
|
||||
t.Fatalf("photo blob ok=%v err=%v", ok, err)
|
||||
}
|
||||
got, err := blobs.Get(ctx, blob.ObjectKey)
|
||||
if err != nil {
|
||||
t.Fatalf("read photo blob: %v", err)
|
||||
}
|
||||
if !bytes.Equal(got, data) {
|
||||
t.Fatalf("photo blob bytes mismatch")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCreateDocumentFromBytesStoresBodyAndAttributes(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
media := newFakeMediaStore()
|
||||
blobs, err := NewLocalFS(t.TempDir())
|
||||
if err != nil {
|
||||
t.Fatalf("NewLocalFS: %v", err)
|
||||
}
|
||||
svc := NewService(media, blobs, 2, WithVideoThumbnailer(nil))
|
||||
data := []byte("inline document body")
|
||||
spec := domain.DocumentSpec{
|
||||
MimeType: "application/pdf",
|
||||
Attributes: []domain.DocumentAttribute{
|
||||
{Kind: domain.DocAttrFilename, FileName: "inline.pdf"},
|
||||
},
|
||||
}
|
||||
|
||||
doc, err := svc.CreateDocumentFromBytes(ctx, data, spec)
|
||||
if err != nil {
|
||||
t.Fatalf("CreateDocumentFromBytes: %v", err)
|
||||
}
|
||||
if doc.ID == 0 || doc.AccessHash == 0 || doc.Size != int64(len(data)) || doc.MimeType != "application/pdf" || len(doc.Attributes) != 1 {
|
||||
t.Fatalf("document = %+v, want stored document body and attributes", doc)
|
||||
}
|
||||
blob, ok, err := media.GetFileBlob(ctx, fmt.Sprintf("doc:%d", doc.ID))
|
||||
if err != nil || !ok {
|
||||
t.Fatalf("document blob ok=%v err=%v", ok, err)
|
||||
}
|
||||
got, err := blobs.Get(ctx, blob.ObjectKey)
|
||||
if err != nil {
|
||||
t.Fatalf("read document blob: %v", err)
|
||||
}
|
||||
if !bytes.Equal(got, data) || blob.MimeType != "application/pdf" {
|
||||
t.Fatalf("document blob mime=%q bytes=%q", blob.MimeType, got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCreateAvatarMarkupGeneratesDownloadableStaticSizes(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
media := newFakeMediaStore()
|
||||
blobs, err := NewLocalFS(t.TempDir())
|
||||
if err != nil {
|
||||
t.Fatalf("NewLocalFS: %v", err)
|
||||
}
|
||||
svc := NewService(media, blobs, 2)
|
||||
|
||||
photo, err := svc.CreateAvatarMarkup(ctx, domain.PhotoSize{
|
||||
Kind: domain.PhotoSizeKindVideoEmojiMarkup,
|
||||
EmojiID: 99,
|
||||
BackgroundColors: []int{0xff3b30, 0x34c759},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CreateAvatarMarkup: %v", err)
|
||||
}
|
||||
if !domain.PhotoHasVideo(photo.Sizes) {
|
||||
t.Fatalf("avatar markup photo sizes = %+v, want video markup", photo.Sizes)
|
||||
}
|
||||
assertDownloadableAvatarSize(t, svc, photo.ID, "a")
|
||||
assertDownloadableAvatarSize(t, svc, photo.ID, "c")
|
||||
}
|
||||
|
||||
// TestCreateAvatarMarkupComposesEmojiThumbIntoStaticSizes 守护两个行为:
|
||||
// 1. 普通彩色 emoji 合成进静态头像时保留原色(不得染白/变黑);
|
||||
// 2. 贴图含非满 alpha 像素(抗锯齿常态)时不得因预乘溢出整片变黑——曾因把
|
||||
// R=G=B=255、A<255 的非法预乘值喂给 draw.Over 溢出回绕,emoji 输出近黑色。
|
||||
func TestCreateAvatarMarkupComposesEmojiThumbIntoStaticSizes(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
media := newFakeMediaStore()
|
||||
blobs, err := NewLocalFS(t.TempDir())
|
||||
if err != nil {
|
||||
t.Fatalf("NewLocalFS: %v", err)
|
||||
}
|
||||
const emojiID = int64(99)
|
||||
if err := media.PutDocument(ctx, domain.Document{
|
||||
ID: emojiID,
|
||||
MimeType: "application/x-tgsticker",
|
||||
Thumbs: []domain.PhotoSize{{
|
||||
Kind: domain.PhotoSizeKindCached,
|
||||
Type: "m",
|
||||
W: 64,
|
||||
H: 64,
|
||||
Bytes: testTransparentThumbPNG(t),
|
||||
}},
|
||||
}); err != nil {
|
||||
t.Fatalf("PutDocument: %v", err)
|
||||
}
|
||||
svc := NewService(media, blobs, 2)
|
||||
|
||||
photo, err := svc.CreateAvatarMarkup(ctx, domain.PhotoSize{
|
||||
Kind: domain.PhotoSizeKindVideoEmojiMarkup,
|
||||
EmojiID: emojiID,
|
||||
BackgroundColors: []int{0x112233, 0x445566},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CreateAvatarMarkup: %v", err)
|
||||
}
|
||||
|
||||
r, g, b, a := avatarStillCenterPixel(t, svc, photo.ID)
|
||||
if a < 250 {
|
||||
t.Fatalf("center pixel alpha=%d, want opaque still", a)
|
||||
}
|
||||
if r < 200 || g > 90 || b > 90 {
|
||||
t.Fatalf("center pixel rgb=(%d,%d,%d), want red emoji color preserved", r, g, b)
|
||||
}
|
||||
}
|
||||
|
||||
// TestCreateAvatarMarkupTintsTextColorEmojiWhite 守护 text_color custom emoji 的
|
||||
// 白色剪影呈现:染色必须写合法预乘值(R=G=B=A),非满 alpha 像素不得溢出变黑。
|
||||
func TestCreateAvatarMarkupTintsTextColorEmojiWhite(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
media := newFakeMediaStore()
|
||||
blobs, err := NewLocalFS(t.TempDir())
|
||||
if err != nil {
|
||||
t.Fatalf("NewLocalFS: %v", err)
|
||||
}
|
||||
const emojiID = int64(120)
|
||||
if err := media.PutDocument(ctx, domain.Document{
|
||||
ID: emojiID,
|
||||
MimeType: "application/x-tgsticker",
|
||||
Attributes: []domain.DocumentAttribute{{
|
||||
Kind: domain.DocAttrCustomEmoji,
|
||||
TextColor: true,
|
||||
}},
|
||||
Thumbs: []domain.PhotoSize{{
|
||||
Kind: domain.PhotoSizeKindCached,
|
||||
Type: "m",
|
||||
W: 64,
|
||||
H: 64,
|
||||
Bytes: testTransparentThumbPNG(t),
|
||||
}},
|
||||
}); err != nil {
|
||||
t.Fatalf("PutDocument: %v", err)
|
||||
}
|
||||
svc := NewService(media, blobs, 2)
|
||||
|
||||
photo, err := svc.CreateAvatarMarkup(ctx, domain.PhotoSize{
|
||||
Kind: domain.PhotoSizeKindVideoEmojiMarkup,
|
||||
EmojiID: emojiID,
|
||||
BackgroundColors: []int{0x112233, 0x445566},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CreateAvatarMarkup: %v", err)
|
||||
}
|
||||
|
||||
r, g, b, _ := avatarStillCenterPixel(t, svc, photo.ID)
|
||||
if r < 230 || g < 230 || b < 230 {
|
||||
t.Fatalf("center pixel rgb=(%d,%d,%d), want white silhouette for text_color emoji", r, g, b)
|
||||
}
|
||||
}
|
||||
|
||||
func avatarStillCenterPixel(t *testing.T, svc *Service, photoID int64) (r, g, b, a uint32) {
|
||||
t.Helper()
|
||||
chunk, found, err := svc.GetFile(context.Background(), domain.FileDownloadRequest{
|
||||
LocationKey: fmt.Sprintf("photo:%d:c", photoID),
|
||||
Offset: 0,
|
||||
Limit: 1 << 20,
|
||||
})
|
||||
if err != nil || !found {
|
||||
t.Fatalf("avatar c blob found=%v err=%v", found, err)
|
||||
}
|
||||
img, _, err := image.Decode(bytes.NewReader(chunk.Bytes))
|
||||
if err != nil {
|
||||
t.Fatalf("decode avatar still: %v", err)
|
||||
}
|
||||
r, g, b, a = img.At(avatarStillSize/2, avatarStillSize/2).RGBA()
|
||||
return r >> 8, g >> 8, b >> 8, a >> 8
|
||||
}
|
||||
|
||||
func TestCreateAvatarVideoMarkupGeneratesDownloadableStaticSizes(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
media := newFakeMediaStore()
|
||||
blobs, err := NewLocalFS(t.TempDir())
|
||||
if err != nil {
|
||||
t.Fatalf("NewLocalFS: %v", err)
|
||||
}
|
||||
svc := NewService(media, blobs, 2)
|
||||
|
||||
if _, err := svc.SaveFilePart(ctx, 10, 400, 0, []byte("fake-profile-video")); err != nil {
|
||||
t.Fatalf("SaveFilePart: %v", err)
|
||||
}
|
||||
photo, err := svc.CreateAvatarVideoMarkupFromUpload(ctx,
|
||||
domain.UploadedFileRef{OwnerUserID: 10, FileID: 400, Parts: 1, Name: "avatar.mp4"},
|
||||
0.25,
|
||||
domain.PhotoSize{
|
||||
Kind: domain.PhotoSizeKindVideoEmojiMarkup,
|
||||
EmojiID: 100,
|
||||
BackgroundColors: []int{0x536dfe, 0x26a69a},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CreateAvatarVideoMarkupFromUpload: %v", err)
|
||||
}
|
||||
assertDownloadableAvatarSize(t, svc, photo.ID, "a")
|
||||
assertDownloadableAvatarSize(t, svc, photo.ID, "c")
|
||||
chunk, found, err := svc.GetFile(ctx, domain.FileDownloadRequest{
|
||||
LocationKey: fmt.Sprintf("photo:%d:u", photo.ID),
|
||||
Offset: 0,
|
||||
Limit: 1024,
|
||||
})
|
||||
if err != nil || !found {
|
||||
t.Fatalf("video avatar blob found=%v err=%v", found, err)
|
||||
}
|
||||
if string(chunk.Bytes) != "fake-profile-video" {
|
||||
t.Fatalf("video avatar bytes = %q", chunk.Bytes)
|
||||
}
|
||||
}
|
||||
|
||||
// TestCreateAvatarVideoMarkupStillUsesVideoFirstFrame 守护动画头像静态尺寸优先取
|
||||
// 上传视频首帧(客户端真实渲染画面),而不是服务端合成的近似 still。
|
||||
func TestCreateAvatarVideoMarkupStillUsesVideoFirstFrame(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
media := newFakeMediaStore()
|
||||
blobs, err := NewLocalFS(t.TempDir())
|
||||
if err != nil {
|
||||
t.Fatalf("NewLocalFS: %v", err)
|
||||
}
|
||||
frame := testJPEG(t, 640, 640)
|
||||
thumbnailer := &fakeVideoThumbnailer{thumb: frame}
|
||||
svc := NewService(media, blobs, 2, WithVideoThumbnailer(thumbnailer))
|
||||
|
||||
if _, err := svc.SaveFilePart(ctx, 10, 500, 0, []byte("fake-profile-video")); err != nil {
|
||||
t.Fatalf("SaveFilePart: %v", err)
|
||||
}
|
||||
photo, err := svc.CreateAvatarVideoMarkupFromUpload(ctx,
|
||||
domain.UploadedFileRef{OwnerUserID: 10, FileID: 500, Parts: 1, Name: "avatar.mp4"},
|
||||
0,
|
||||
domain.PhotoSize{
|
||||
Kind: domain.PhotoSizeKindVideoEmojiMarkup,
|
||||
EmojiID: 77,
|
||||
BackgroundColors: []int{0x112233},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CreateAvatarVideoMarkupFromUpload: %v", err)
|
||||
}
|
||||
if thumbnailer.calls != 1 {
|
||||
t.Fatalf("thumbnailer calls = %d, want 1", thumbnailer.calls)
|
||||
}
|
||||
chunk, found, err := svc.GetFile(ctx, domain.FileDownloadRequest{
|
||||
LocationKey: fmt.Sprintf("photo:%d:a", photo.ID),
|
||||
Offset: 0,
|
||||
Limit: 1 << 20,
|
||||
})
|
||||
if err != nil || !found {
|
||||
t.Fatalf("avatar a blob found=%v err=%v", found, err)
|
||||
}
|
||||
if !bytes.Equal(chunk.Bytes, frame) {
|
||||
t.Fatalf("avatar still bytes != extracted first frame (got %d bytes, want %d)", len(chunk.Bytes), len(frame))
|
||||
}
|
||||
if chunk.MimeType != "image/jpeg" {
|
||||
t.Fatalf("avatar still mime = %q, want image/jpeg from extracted frame", chunk.MimeType)
|
||||
}
|
||||
}
|
||||
|
||||
func assertDownloadableAvatarSize(t *testing.T, svc *Service, photoID int64, sizeType string) {
|
||||
t.Helper()
|
||||
chunk, found, err := svc.GetFile(context.Background(), domain.FileDownloadRequest{
|
||||
LocationKey: fmt.Sprintf("photo:%d:%s", photoID, sizeType),
|
||||
Offset: 0,
|
||||
Limit: 1 << 20,
|
||||
})
|
||||
if err != nil || !found {
|
||||
t.Fatalf("avatar %s blob found=%v err=%v", sizeType, found, err)
|
||||
}
|
||||
if len(chunk.Bytes) == 0 || chunk.MimeType != "image/png" {
|
||||
t.Fatalf("avatar %s chunk mime=%q bytes=%d, want image/png bytes", sizeType, chunk.MimeType, len(chunk.Bytes))
|
||||
}
|
||||
}
|
||||
|
||||
// testTransparentThumbPNG 构造红色方块贴图:周边透明、中心 alpha=250(模拟抗锯齿
|
||||
// 的非满 alpha),用于守护预乘溢出回归——溢出代码会把 alpha≠255 的像素整片渲染成黑。
|
||||
func testTransparentThumbPNG(t *testing.T) []byte {
|
||||
t.Helper()
|
||||
img := image.NewNRGBA(image.Rect(0, 0, 64, 64))
|
||||
for y := 8; y < 56; y++ {
|
||||
for x := 8; x < 56; x++ {
|
||||
img.SetNRGBA(x, y, color.NRGBA{R: 240, G: 30, B: 30, A: 250})
|
||||
}
|
||||
}
|
||||
var buf bytes.Buffer
|
||||
if err := png.Encode(&buf, img); err != nil {
|
||||
t.Fatalf("encode test thumb: %v", err)
|
||||
}
|
||||
return buf.Bytes()
|
||||
}
|
||||
|
||||
type fakeVideoThumbnailer struct {
|
||||
calls int
|
||||
thumb []byte
|
||||
err error
|
||||
}
|
||||
|
||||
func (f *fakeVideoThumbnailer) Extract(context.Context, []byte, string) ([]byte, error) {
|
||||
f.calls++
|
||||
if f.err != nil {
|
||||
return nil, f.err
|
||||
}
|
||||
return append([]byte(nil), f.thumb...), nil
|
||||
}
|
||||
|
||||
func testJPEG(t *testing.T, w, h int) []byte {
|
||||
t.Helper()
|
||||
img := image.NewRGBA(image.Rect(0, 0, w, h))
|
||||
for y := 0; y < h; y++ {
|
||||
for x := 0; x < w; x++ {
|
||||
img.Set(x, y, color.RGBA{R: uint8(40 + x), G: uint8(80 + y), B: 120, A: 255})
|
||||
}
|
||||
}
|
||||
var buf bytes.Buffer
|
||||
if err := jpeg.Encode(&buf, img, &jpeg.Options{Quality: 90}); err != nil {
|
||||
t.Fatalf("encode jpeg: %v", err)
|
||||
}
|
||||
return buf.Bytes()
|
||||
}
|
||||
|
|
@ -1,10 +1,14 @@
|
|||
package files
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"image"
|
||||
"image/color"
|
||||
"image/png"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"regexp"
|
||||
|
|
@ -13,6 +17,8 @@ import (
|
|||
"strings"
|
||||
"time"
|
||||
|
||||
"go.uber.org/zap"
|
||||
|
||||
"telesrv/internal/domain"
|
||||
)
|
||||
|
||||
|
|
@ -25,6 +31,7 @@ import (
|
|||
type SeedStats struct {
|
||||
Reactions int
|
||||
StickerSets int
|
||||
Effects int
|
||||
Documents int
|
||||
Blobs int
|
||||
Skipped bool
|
||||
|
|
@ -38,11 +45,19 @@ func (s *Service) SeedMedia(ctx context.Context, root string, maxRegularSets int
|
|||
return stats, nil
|
||||
}
|
||||
if _, err := os.Stat(root); err != nil {
|
||||
// 目录不存在:跳过而非失败(开发机可能未放资源)。
|
||||
// 目录不存在:跳过而非失败(开发机可能未放资源)。但显式 WARN——否则「配置的 seed 目录
|
||||
// 不存在 → 静默不导入贴纸/reaction」会被埋没(DB 已有旧数据时尤其隐蔽,表现为客户端反复
|
||||
// 拉取未 seed 的集)。配置 TELESRV_STICKER_SEED_DIR 指向真实导出目录即可。
|
||||
if s.log != nil {
|
||||
s.log.Warn("sticker/reaction seed 目录不存在,跳过媒体种子导入(配置 TELESRV_STICKER_SEED_DIR)",
|
||||
zap.String("dir", root), zap.Error(err))
|
||||
}
|
||||
stats.Skipped = true
|
||||
return stats, nil
|
||||
}
|
||||
|
||||
phaseStarted := time.Now()
|
||||
phaseBefore := stats
|
||||
// reactions
|
||||
if n, err := s.media.CountAvailableReactions(ctx); err != nil {
|
||||
return stats, err
|
||||
|
|
@ -57,28 +72,58 @@ func (s *Service) SeedMedia(ctx context.Context, root string, maxRegularSets int
|
|||
return stats, fmt.Errorf("repair reactions: %w", err)
|
||||
}
|
||||
}
|
||||
s.logSeedPhase("reactions", phaseStarted, phaseBefore, stats)
|
||||
|
||||
// sticker sets(default 系统集 + 常规集)
|
||||
phaseStarted = time.Now()
|
||||
phaseBefore = stats
|
||||
// sticker sets(default 系统集 + 常规集 + emoji 集):始终扫描导出目录以**增量**拾取
|
||||
// 新增的 set 目录——importStickerSetDir 跳过内容(hash)未变的已有集,只导入新集/变更集。
|
||||
// 这样向已部署(非空 store)的 data/sticker-seed 丢新集后重启即可生效,无需清库重 seed。
|
||||
// 仅当检测到旧版缩略图/可渲染预览元数据缺失时 force=true 全量重导修复。
|
||||
forceSticker := false
|
||||
if n, err := s.media.CountStickerSets(ctx); err != nil {
|
||||
return stats, err
|
||||
} else if n == 0 {
|
||||
if err := s.seedStickerSets(ctx, root, maxRegularSets, &stats); err != nil {
|
||||
return stats, fmt.Errorf("seed sticker sets: %w", err)
|
||||
}
|
||||
} else if stale, err := s.stickerSetDocumentThumbsNeedInlineCache(ctx); err != nil {
|
||||
return stats, err
|
||||
} else if stale {
|
||||
if err := s.seedStickerSets(ctx, root, maxRegularSets, &stats); err != nil {
|
||||
return stats, fmt.Errorf("repair sticker set thumbs: %w", err)
|
||||
} else if n > 0 {
|
||||
stale, err := s.stickerSetDocumentsNeedSeedRepair(ctx)
|
||||
if err != nil {
|
||||
return stats, err
|
||||
}
|
||||
forceSticker = stale
|
||||
}
|
||||
if err := s.seedStickerSets(ctx, root, maxRegularSets, forceSticker, &stats); err != nil {
|
||||
return stats, fmt.Errorf("seed sticker sets: %w", err)
|
||||
}
|
||||
s.logSeedPhase("sticker_sets", phaseStarted, phaseBefore, stats)
|
||||
|
||||
if stats.Reactions == 0 && stats.StickerSets == 0 {
|
||||
phaseStarted = time.Now()
|
||||
phaseBefore = stats
|
||||
// 消息发送特效:全局静态目录,每次启动重建内存 s.effects(文档导入幂等)。
|
||||
if err := s.seedEffects(ctx, root, &stats); err != nil {
|
||||
return stats, fmt.Errorf("seed effects: %w", err)
|
||||
}
|
||||
s.logSeedPhase("effects", phaseStarted, phaseBefore, stats)
|
||||
|
||||
if stats.Reactions == 0 && stats.StickerSets == 0 && stats.Effects == 0 {
|
||||
stats.Skipped = true
|
||||
}
|
||||
return stats, nil
|
||||
}
|
||||
|
||||
func (s *Service) logSeedPhase(phase string, started time.Time, before, after SeedStats) {
|
||||
if s.log == nil {
|
||||
return
|
||||
}
|
||||
s.log.Info("媒体种子阶段完成",
|
||||
zap.String("phase", phase),
|
||||
zap.Duration("elapsed", time.Since(started)),
|
||||
zap.Int("reactions", after.Reactions-before.Reactions),
|
||||
zap.Int("sticker_sets", after.StickerSets-before.StickerSets),
|
||||
zap.Int("effects", after.Effects-before.Effects),
|
||||
zap.Int("documents", after.Documents-before.Documents),
|
||||
zap.Int("blobs", after.Blobs-before.Blobs),
|
||||
)
|
||||
}
|
||||
|
||||
// ---- reactions ----
|
||||
|
||||
func (s *Service) seedReactions(ctx context.Context, root string, stats *SeedStats) error {
|
||||
|
|
@ -175,7 +220,7 @@ func (s *Service) availableReactionSeedNeedsRepair(ctx context.Context) (bool, e
|
|||
|
||||
// ---- sticker sets ----
|
||||
|
||||
func (s *Service) seedStickerSets(ctx context.Context, root string, maxRegular int, stats *SeedStats) error {
|
||||
func (s *Service) seedStickerSets(ctx context.Context, root string, maxRegular int, force bool, stats *SeedStats) error {
|
||||
// default 系统集:目录名 → system_key。
|
||||
defaultDir := filepath.Join(root, "telegram_default_stickers_export")
|
||||
order := 0
|
||||
|
|
@ -190,13 +235,33 @@ func (s *Service) seedStickerSets(ctx context.Context, root string, maxRegular i
|
|||
for _, name := range names {
|
||||
systemKey := systemKeyForDefaultSet(name)
|
||||
setDir := filepath.Join(defaultDir, name)
|
||||
if err := s.importStickerSetDir(ctx, setDir, systemKey, order, stats); err != nil {
|
||||
if err := s.importStickerSetDir(ctx, setDir, systemKey, order, force, stats); err != nil {
|
||||
return fmt.Errorf("import default set %s: %w", name, err)
|
||||
}
|
||||
order++
|
||||
}
|
||||
}
|
||||
|
||||
// custom-emoji 集(telegram_emoji_export/<set>/):不受 maxRegular 限制,按 set_info 的
|
||||
// emojis 标志归入 StickerSetKindEmoji(getEmojiStickers/getFeaturedEmojiStickers 下发)。
|
||||
emojiDir := filepath.Join(root, "telegram_emoji_export")
|
||||
if entries, err := os.ReadDir(emojiDir); err == nil {
|
||||
names := make([]string, 0, len(entries))
|
||||
for _, e := range entries {
|
||||
if e.IsDir() {
|
||||
names = append(names, e.Name())
|
||||
}
|
||||
}
|
||||
sort.Strings(names)
|
||||
for _, name := range names {
|
||||
setDir := filepath.Join(emojiDir, name)
|
||||
if err := s.importStickerSetDir(ctx, setDir, "", order, force, stats); err != nil {
|
||||
return fmt.Errorf("import emoji set %s: %w", name, err)
|
||||
}
|
||||
order++
|
||||
}
|
||||
}
|
||||
|
||||
// 常规贴纸集。
|
||||
regularDir := filepath.Join(root, "telegram_stickers_export")
|
||||
if entries, err := os.ReadDir(regularDir); err == nil {
|
||||
|
|
@ -213,7 +278,7 @@ func (s *Service) seedStickerSets(ctx context.Context, root string, maxRegular i
|
|||
break
|
||||
}
|
||||
setDir := filepath.Join(regularDir, name)
|
||||
if err := s.importStickerSetDir(ctx, setDir, "", order, stats); err != nil {
|
||||
if err := s.importStickerSetDir(ctx, setDir, "", order, force, stats); err != nil {
|
||||
return fmt.Errorf("import sticker set %s: %w", name, err)
|
||||
}
|
||||
order++
|
||||
|
|
@ -223,7 +288,7 @@ func (s *Service) seedStickerSets(ctx context.Context, root string, maxRegular i
|
|||
return nil
|
||||
}
|
||||
|
||||
func (s *Service) importStickerSetDir(ctx context.Context, setDir, systemKey string, order int, stats *SeedStats) error {
|
||||
func (s *Service) importStickerSetDir(ctx context.Context, setDir, systemKey string, order int, force bool, stats *SeedStats) error {
|
||||
infoPath := filepath.Join(setDir, "set_info.json")
|
||||
raw, err := os.ReadFile(infoPath)
|
||||
if err != nil {
|
||||
|
|
@ -239,6 +304,16 @@ func (s *Service) importStickerSetDir(ctx context.Context, setDir, systemKey str
|
|||
if sj.ID == 0 {
|
||||
return nil
|
||||
}
|
||||
// 增量 seed:集已存在且内容 hash 未变则跳过(不重读文档/重传 blob)。force 时强制重导
|
||||
// (缩略图内联缓存修复路径)。这让 seedStickerSets 可在非空 store 上每次启动安全重扫,
|
||||
// 仅导入新增/变更集。
|
||||
if !force {
|
||||
if existing, found, err := s.media.GetStickerSetByID(ctx, sj.ID); err != nil {
|
||||
return err
|
||||
} else if found && existing.Hash == sj.Hash {
|
||||
return nil
|
||||
}
|
||||
}
|
||||
stickersDir := filepath.Join(setDir, "stickers")
|
||||
index, err := scanSeedDir(stickersDir)
|
||||
if err != nil {
|
||||
|
|
@ -382,6 +457,10 @@ func (s *Service) importDocument(ctx context.Context, dj seedDocumentJSON, binDi
|
|||
}
|
||||
doc.Thumbs = thumbs
|
||||
|
||||
if err := s.ensureTGStickerPreviewThumb(ctx, &doc, stats); err != nil {
|
||||
return domain.Document{}, err
|
||||
}
|
||||
|
||||
if err := s.media.PutDocument(ctx, doc); err != nil {
|
||||
return domain.Document{}, err
|
||||
}
|
||||
|
|
@ -406,6 +485,7 @@ var seedTrailingDigits = regexp.MustCompile(`(\d{6,})`)
|
|||
var seedThumbMarker = regexp.MustCompile(`_thumb\d+_`)
|
||||
|
||||
const seedInlineCachedDocumentThumbMaxBytes = 32 * 1024
|
||||
const seedSyntheticDocumentThumbType = "m"
|
||||
|
||||
// Exported Telegram resources keep their original id in filenames/JSON, but
|
||||
// telesrv owns the document catalog it serves. Imported high source ids are
|
||||
|
|
@ -413,6 +493,7 @@ const seedInlineCachedDocumentThumbMaxBytes = 32 * 1024
|
|||
const seedExternalDocumentIDOffset int64 = 4_000_000_000_000_000_000
|
||||
|
||||
var seedThumbType = regexp.MustCompile(`PhotoSize_type([a-z])`)
|
||||
var seedSyntheticTGStickerPreviewThumbPNG = makeSeedSyntheticTGStickerPreviewThumbPNG()
|
||||
|
||||
func seedDocumentStorageID(sourceID int64) int64 {
|
||||
if sourceID <= 0 {
|
||||
|
|
@ -559,17 +640,6 @@ func seedDocumentAttributes(attrs []seedAttrJSON) []domain.DocumentAttribute {
|
|||
return out
|
||||
}
|
||||
|
||||
func seedPhotoSizes(thumbs []seedThumbJSON) []domain.PhotoSize {
|
||||
out := make([]domain.PhotoSize, 0, len(thumbs))
|
||||
for _, t := range thumbs {
|
||||
ps, _ := seedPhotoSize(t)
|
||||
if ps.Kind != "" {
|
||||
out = append(out, ps)
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func seedStickerSetPhotoSizes(thumbs []seedThumbJSON) []domain.PhotoSize {
|
||||
out := make([]domain.PhotoSize, 0, len(thumbs))
|
||||
for _, t := range thumbs {
|
||||
|
|
@ -612,6 +682,63 @@ func seedInlineCachedDocumentThumb(ps domain.PhotoSize, data []byte) domain.Phot
|
|||
return ps
|
||||
}
|
||||
|
||||
func (s *Service) ensureTGStickerPreviewThumb(ctx context.Context, doc *domain.Document, stats *SeedStats) error {
|
||||
if !seedDocumentNeedsSyntheticTGStickerPreviewThumb(*doc) {
|
||||
return nil
|
||||
}
|
||||
if s.blobs == nil {
|
||||
return fmt.Errorf("blob backend not configured for synthetic sticker preview thumb")
|
||||
}
|
||||
data := seedSyntheticTGStickerPreviewThumbPNG
|
||||
objectKey, err := s.blobs.Put(ctx, data)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := s.media.PutFileBlob(ctx, domain.FileBlob{
|
||||
LocationKey: fmt.Sprintf("doc:%d:%s", doc.ID, seedSyntheticDocumentThumbType),
|
||||
Backend: domain.MediaBackend(s.blobs.Name()),
|
||||
ObjectKey: objectKey,
|
||||
Size: int64(len(data)),
|
||||
MimeType: "image/png",
|
||||
}); err != nil {
|
||||
return err
|
||||
}
|
||||
doc.Thumbs = append(doc.Thumbs, domain.PhotoSize{
|
||||
Kind: domain.PhotoSizeKindCached,
|
||||
Type: seedSyntheticDocumentThumbType,
|
||||
W: 1,
|
||||
H: 1,
|
||||
Bytes: append([]byte(nil), data...),
|
||||
})
|
||||
s.prewarmSmallBlob(objectKey, data)
|
||||
stats.Blobs++
|
||||
return nil
|
||||
}
|
||||
|
||||
func seedDocumentNeedsSyntheticTGStickerPreviewThumb(doc domain.Document) bool {
|
||||
if doc.MimeType != "application/x-tgsticker" || len(doc.Thumbs) > 0 {
|
||||
return false
|
||||
}
|
||||
return seedDocumentHasAttribute(doc.Attributes, domain.DocAttrCustomEmoji)
|
||||
}
|
||||
|
||||
func seedDocumentHasAttribute(attrs []domain.DocumentAttribute, kind domain.DocumentAttributeKind) bool {
|
||||
for _, attr := range attrs {
|
||||
if attr.Kind == kind {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func makeSeedSyntheticTGStickerPreviewThumbPNG() []byte {
|
||||
var buf bytes.Buffer
|
||||
img := image.NewNRGBA(image.Rect(0, 0, 1, 1))
|
||||
img.Set(0, 0, color.NRGBA{})
|
||||
_ = png.Encode(&buf, img)
|
||||
return buf.Bytes()
|
||||
}
|
||||
|
||||
func seedThumbMimeType(data []byte) string {
|
||||
switch {
|
||||
case len(data) >= 12 && data[0] == 'R' && data[1] == 'I' && data[2] == 'F' && data[3] == 'F' &&
|
||||
|
|
@ -628,7 +755,7 @@ func seedThumbMimeType(data []byte) string {
|
|||
}
|
||||
}
|
||||
|
||||
func (s *Service) stickerSetDocumentThumbsNeedInlineCache(ctx context.Context) (bool, error) {
|
||||
func (s *Service) stickerSetDocumentsNeedSeedRepair(ctx context.Context) (bool, error) {
|
||||
var ids []int64
|
||||
for _, kind := range []domain.StickerSetKind{
|
||||
domain.StickerSetKindStickers,
|
||||
|
|
@ -644,10 +771,14 @@ func (s *Service) stickerSetDocumentThumbsNeedInlineCache(ctx context.Context) (
|
|||
ids = append(ids, set.DocumentIDs...)
|
||||
}
|
||||
}
|
||||
return s.documentsNeedInlineCachedThumbs(ctx, ids)
|
||||
return s.documentsNeedSeedRepair(ctx, ids)
|
||||
}
|
||||
|
||||
func (s *Service) documentsNeedInlineCachedThumbs(ctx context.Context, ids []int64) (bool, error) {
|
||||
return s.documentsNeedSeedRepair(ctx, ids)
|
||||
}
|
||||
|
||||
func (s *Service) documentsNeedSeedRepair(ctx context.Context, ids []int64) (bool, error) {
|
||||
if len(ids) == 0 {
|
||||
return false, nil
|
||||
}
|
||||
|
|
@ -668,6 +799,9 @@ func (s *Service) documentsNeedInlineCachedThumbs(ctx context.Context, ids []int
|
|||
return false, err
|
||||
}
|
||||
for _, doc := range docs {
|
||||
if seedDocumentNeedsSyntheticTGStickerPreviewThumb(doc) {
|
||||
return true, nil
|
||||
}
|
||||
for _, thumb := range doc.Thumbs {
|
||||
if thumb.Kind == domain.PhotoSizeKindDefault && thumb.Size > 0 && thumb.Size <= seedInlineCachedDocumentThumbMaxBytes {
|
||||
return true, nil
|
||||
|
|
|
|||
172
internal/app/files/seed_state.go
Normal file
172
internal/app/files/seed_state.go
Normal file
|
|
@ -0,0 +1,172 @@
|
|||
package files
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"fmt"
|
||||
"hash"
|
||||
"io"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sort"
|
||||
|
||||
"telesrv/internal/domain"
|
||||
)
|
||||
|
||||
const (
|
||||
seedEffectsStateKey = "files.effects"
|
||||
seedEffectsStateVersion = "effects-v2"
|
||||
seedAppearanceStateKey = "files.appearance"
|
||||
seedAppearanceStateVersion = "appearance-v1"
|
||||
)
|
||||
|
||||
func (s *Service) seedStateMatches(ctx context.Context, key, want string) (bool, error) {
|
||||
if want == "" {
|
||||
return false, nil
|
||||
}
|
||||
got, found, err := s.media.GetSeedState(ctx, key)
|
||||
if err != nil || !found {
|
||||
return false, err
|
||||
}
|
||||
return got == want, nil
|
||||
}
|
||||
|
||||
func (s *Service) putSeedState(ctx context.Context, key, hash string) error {
|
||||
if hash == "" {
|
||||
return nil
|
||||
}
|
||||
return s.media.PutSeedState(ctx, key, hash)
|
||||
}
|
||||
|
||||
func seedStateHash(write func(hash.Hash) error) (string, error) {
|
||||
h := sha256.New()
|
||||
if err := write(h); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return hex.EncodeToString(h.Sum(nil)), nil
|
||||
}
|
||||
|
||||
func writeSeedStateHeader(h io.Writer, version string, dc int) {
|
||||
_, _ = fmt.Fprintf(h, "version=%s\ndc=%d\n", version, dc)
|
||||
}
|
||||
|
||||
func writeSeedDirFingerprint(h io.Writer, dir string) error {
|
||||
entries, err := os.ReadDir(dir)
|
||||
if err != nil {
|
||||
if os.IsNotExist(err) {
|
||||
return nil
|
||||
}
|
||||
return err
|
||||
}
|
||||
names := make([]string, 0, len(entries))
|
||||
byName := make(map[string]os.DirEntry, len(entries))
|
||||
for _, entry := range entries {
|
||||
if entry.IsDir() {
|
||||
continue
|
||||
}
|
||||
names = append(names, entry.Name())
|
||||
byName[entry.Name()] = entry
|
||||
}
|
||||
sort.Strings(names)
|
||||
for _, name := range names {
|
||||
info, err := byName[name].Info()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
rel := filepath.ToSlash(name)
|
||||
_, _ = fmt.Fprintf(h, "file=%s\x00size=%d\x00mtime=%d\n", rel, info.Size(), info.ModTime().UnixNano())
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func seedDocumentJSONLocationKeys(dj seedDocumentJSON, index seedDirIndex) []string {
|
||||
storageID := seedDocumentStorageID(dj.ID)
|
||||
if storageID == 0 {
|
||||
return nil
|
||||
}
|
||||
keys := make([]string, 0, 1+len(dj.Thumbs))
|
||||
if _, ok := index.main[dj.ID]; ok {
|
||||
keys = append(keys, fmt.Sprintf("doc:%d", storageID))
|
||||
}
|
||||
for _, tj := range dj.Thumbs {
|
||||
ps, downloadable := seedPhotoSize(tj)
|
||||
if !downloadable || ps.Type == "" {
|
||||
continue
|
||||
}
|
||||
if _, ok := index.thumb[dj.ID][ps.Type]; ok {
|
||||
keys = append(keys, fmt.Sprintf("doc:%d:%s", storageID, ps.Type))
|
||||
}
|
||||
}
|
||||
if seedDocumentJSONNeedsSyntheticTGStickerPreviewThumb(dj) {
|
||||
keys = append(keys, fmt.Sprintf("doc:%d:%s", storageID, seedSyntheticDocumentThumbType))
|
||||
}
|
||||
return keys
|
||||
}
|
||||
|
||||
func seedDocumentJSONNeedsSyntheticTGStickerPreviewThumb(dj seedDocumentJSON) bool {
|
||||
if dj.MimeType != "application/x-tgsticker" || len(dj.Thumbs) > 0 {
|
||||
return false
|
||||
}
|
||||
return seedDocumentHasAttribute(seedDocumentAttributes(dj.Attributes), domain.DocAttrCustomEmoji)
|
||||
}
|
||||
|
||||
func (s *Service) seedDocumentJSONsReady(ctx context.Context, docs []seedDocumentJSON, index seedDirIndex) (bool, error) {
|
||||
expected := make(map[int64]seedDocumentJSON, len(docs))
|
||||
ids := make([]int64, 0, len(docs))
|
||||
locationKeys := make([]string, 0, len(docs))
|
||||
seenLocationKeys := make(map[string]struct{}, len(docs))
|
||||
for _, dj := range docs {
|
||||
storageID := seedDocumentStorageID(dj.ID)
|
||||
if storageID == 0 {
|
||||
continue
|
||||
}
|
||||
if _, ok := expected[storageID]; !ok {
|
||||
expected[storageID] = dj
|
||||
ids = append(ids, storageID)
|
||||
}
|
||||
for _, key := range seedDocumentJSONLocationKeys(dj, index) {
|
||||
if _, ok := seenLocationKeys[key]; ok {
|
||||
continue
|
||||
}
|
||||
seenLocationKeys[key] = struct{}{}
|
||||
locationKeys = append(locationKeys, key)
|
||||
}
|
||||
}
|
||||
if len(ids) == 0 {
|
||||
return true, nil
|
||||
}
|
||||
stored, err := s.media.GetDocuments(ctx, ids)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
if len(stored) < len(expected) {
|
||||
return false, nil
|
||||
}
|
||||
for _, doc := range stored {
|
||||
dj, ok := expected[doc.ID]
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
if doc.DCID != s.dc || doc.MimeType != dj.MimeType || doc.Size != dj.Size {
|
||||
return false, nil
|
||||
}
|
||||
delete(expected, doc.ID)
|
||||
}
|
||||
if len(expected) > 0 {
|
||||
return false, nil
|
||||
}
|
||||
if len(locationKeys) == 0 {
|
||||
return true, nil
|
||||
}
|
||||
blobs, err := s.media.GetFileBlobs(ctx, locationKeys)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
for _, key := range locationKeys {
|
||||
if _, ok := blobs[key]; !ok {
|
||||
return false, nil
|
||||
}
|
||||
}
|
||||
return true, nil
|
||||
}
|
||||
|
|
@ -2,10 +2,13 @@ package files
|
|||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sort"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"telesrv/internal/domain"
|
||||
)
|
||||
|
|
@ -19,23 +22,96 @@ type fakeMediaStore struct {
|
|||
sets map[int64]domain.StickerSet
|
||||
reactions []domain.AvailableReaction
|
||||
parts map[string][]domain.UploadPart
|
||||
webPages map[int64]domain.MessageWebPage
|
||||
seedState map[string]string
|
||||
}
|
||||
|
||||
func newFakeMediaStore() *fakeMediaStore {
|
||||
return &fakeMediaStore{
|
||||
blobs: map[string]domain.FileBlob{},
|
||||
docs: map[int64]domain.Document{},
|
||||
photos: map[int64]domain.Photo{},
|
||||
sets: map[int64]domain.StickerSet{},
|
||||
parts: map[string][]domain.UploadPart{},
|
||||
blobs: map[string]domain.FileBlob{},
|
||||
docs: map[int64]domain.Document{},
|
||||
photos: map[int64]domain.Photo{},
|
||||
sets: map[int64]domain.StickerSet{},
|
||||
parts: map[string][]domain.UploadPart{},
|
||||
seedState: map[string]string{},
|
||||
}
|
||||
}
|
||||
|
||||
func (f *fakeMediaStore) SaveFilePart(_ context.Context, _ domain.UploadPart) error { return nil }
|
||||
func (f *fakeMediaStore) LoadFileParts(_ context.Context, _, _ int64) ([]domain.UploadPart, error) {
|
||||
func (f *fakeMediaStore) SaveFilePart(_ context.Context, part domain.UploadPart) error {
|
||||
f.mu.Lock()
|
||||
defer f.mu.Unlock()
|
||||
key := fakeUploadPartKey(part.OwnerUserID, part.FileID)
|
||||
part.SHA256 = append([]byte(nil), part.SHA256...)
|
||||
parts := f.parts[key]
|
||||
for i := range parts {
|
||||
if parts[i].Part == part.Part {
|
||||
parts[i] = part
|
||||
f.parts[key] = parts
|
||||
return nil
|
||||
}
|
||||
}
|
||||
f.parts[key] = append(parts, part)
|
||||
return nil
|
||||
}
|
||||
func (f *fakeMediaStore) UploadPartUsage(_ context.Context, ownerUserID int64) (domain.UploadPartUsage, error) {
|
||||
f.mu.Lock()
|
||||
defer f.mu.Unlock()
|
||||
var usage domain.UploadPartUsage
|
||||
files := map[int64]struct{}{}
|
||||
for _, parts := range f.parts {
|
||||
for _, p := range parts {
|
||||
if p.OwnerUserID != ownerUserID {
|
||||
continue
|
||||
}
|
||||
usage.Bytes += p.Size
|
||||
usage.Parts++
|
||||
files[p.FileID] = struct{}{}
|
||||
}
|
||||
}
|
||||
usage.Files = len(files)
|
||||
return usage, nil
|
||||
}
|
||||
func (f *fakeMediaStore) UploadPartSlot(_ context.Context, ownerUserID, fileID int64, part int) (domain.UploadPartSlot, error) {
|
||||
f.mu.Lock()
|
||||
defer f.mu.Unlock()
|
||||
parts := f.parts[fakeUploadPartKey(ownerUserID, fileID)]
|
||||
slot := domain.UploadPartSlot{FileParts: len(parts)}
|
||||
for _, p := range parts {
|
||||
if p.Part == part {
|
||||
slot.ExistingBytes = p.Size
|
||||
slot.ObjectKey = p.ObjectKey
|
||||
slot.Found = true
|
||||
break
|
||||
}
|
||||
}
|
||||
return slot, nil
|
||||
}
|
||||
func (f *fakeMediaStore) LoadFileParts(_ context.Context, ownerUserID, fileID int64) ([]domain.UploadPart, error) {
|
||||
f.mu.Lock()
|
||||
defer f.mu.Unlock()
|
||||
parts := append([]domain.UploadPart(nil), f.parts[fakeUploadPartKey(ownerUserID, fileID)]...)
|
||||
sort.Slice(parts, func(i, j int) bool { return parts[i].Part < parts[j].Part })
|
||||
return parts, nil
|
||||
}
|
||||
func (f *fakeMediaStore) DeleteFileParts(_ context.Context, ownerUserID, fileID int64) ([]string, error) {
|
||||
f.mu.Lock()
|
||||
defer f.mu.Unlock()
|
||||
key := fakeUploadPartKey(ownerUserID, fileID)
|
||||
parts := f.parts[key]
|
||||
keys := make([]string, 0, len(parts))
|
||||
for _, p := range parts {
|
||||
keys = append(keys, p.ObjectKey)
|
||||
}
|
||||
delete(f.parts, key)
|
||||
return keys, nil
|
||||
}
|
||||
func (f *fakeMediaStore) DeleteExpiredUploadParts(_ context.Context, _ time.Time, _ int) ([]string, error) {
|
||||
return nil, nil
|
||||
}
|
||||
func (f *fakeMediaStore) DeleteFileParts(_ context.Context, _, _ int64) error { return nil }
|
||||
|
||||
func fakeUploadPartKey(ownerUserID, fileID int64) string {
|
||||
return fmt.Sprintf("%d:%d", ownerUserID, fileID)
|
||||
}
|
||||
|
||||
func (f *fakeMediaStore) PutFileBlob(_ context.Context, blob domain.FileBlob) error {
|
||||
f.mu.Lock()
|
||||
|
|
@ -50,6 +126,32 @@ func (f *fakeMediaStore) GetFileBlob(_ context.Context, key string) (domain.File
|
|||
return b, ok, nil
|
||||
}
|
||||
|
||||
func (f *fakeMediaStore) GetFileBlobs(_ context.Context, keys []string) (map[string]domain.FileBlob, error) {
|
||||
f.mu.Lock()
|
||||
defer f.mu.Unlock()
|
||||
out := make(map[string]domain.FileBlob, len(keys))
|
||||
for _, key := range keys {
|
||||
if b, ok := f.blobs[key]; ok {
|
||||
out[key] = b
|
||||
}
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
func (f *fakeMediaStore) GetSeedState(_ context.Context, key string) (string, bool, error) {
|
||||
f.mu.Lock()
|
||||
defer f.mu.Unlock()
|
||||
hash, ok := f.seedState[key]
|
||||
return hash, ok, nil
|
||||
}
|
||||
|
||||
func (f *fakeMediaStore) PutSeedState(_ context.Context, key, hash string) error {
|
||||
f.mu.Lock()
|
||||
defer f.mu.Unlock()
|
||||
f.seedState[key] = hash
|
||||
return nil
|
||||
}
|
||||
|
||||
func (f *fakeMediaStore) PutDocument(_ context.Context, doc domain.Document) error {
|
||||
f.mu.Lock()
|
||||
defer f.mu.Unlock()
|
||||
|
|
@ -85,6 +187,21 @@ func (f *fakeMediaStore) GetPhoto(_ context.Context, id int64) (domain.Photo, bo
|
|||
p, ok := f.photos[id]
|
||||
return p, ok, nil
|
||||
}
|
||||
func (f *fakeMediaStore) PutWebPage(_ context.Context, urlHash int64, page domain.MessageWebPage, _ int) error {
|
||||
f.mu.Lock()
|
||||
defer f.mu.Unlock()
|
||||
if f.webPages == nil {
|
||||
f.webPages = map[int64]domain.MessageWebPage{}
|
||||
}
|
||||
f.webPages[urlHash] = page
|
||||
return nil
|
||||
}
|
||||
func (f *fakeMediaStore) GetWebPageByURLHash(_ context.Context, urlHash int64) (domain.MessageWebPage, int, bool, error) {
|
||||
f.mu.Lock()
|
||||
defer f.mu.Unlock()
|
||||
p, ok := f.webPages[urlHash]
|
||||
return p, 0, ok, nil
|
||||
}
|
||||
|
||||
func (f *fakeMediaStore) PutStickerSet(_ context.Context, set domain.StickerSet) error {
|
||||
f.mu.Lock()
|
||||
|
|
@ -156,15 +273,9 @@ func (f *fakeMediaStore) CountAvailableReactions(_ context.Context) (int, error)
|
|||
defer f.mu.Unlock()
|
||||
return len(f.reactions), nil
|
||||
}
|
||||
func (f *fakeMediaStore) AddProfilePhoto(_ context.Context, _ domain.PeerType, _, _ int64, _ int) error {
|
||||
return nil
|
||||
}
|
||||
func (f *fakeMediaStore) AddProfilePhotoKind(_ context.Context, _ domain.PeerType, _ int64, _ domain.ProfilePhotoKind, _ int64, _ int) error {
|
||||
return nil
|
||||
}
|
||||
func (f *fakeMediaStore) CurrentProfilePhoto(_ context.Context, _ domain.PeerType, _ int64) (int64, bool, error) {
|
||||
return 0, false, nil
|
||||
}
|
||||
func (f *fakeMediaStore) CurrentProfilePhotoKind(_ context.Context, _ domain.PeerType, _ int64, _ domain.ProfilePhotoKind) (int64, bool, error) {
|
||||
return 0, false, nil
|
||||
}
|
||||
|
|
@ -174,10 +285,10 @@ func (f *fakeMediaStore) CurrentProfilePhotos(_ context.Context, _ domain.PeerTy
|
|||
func (f *fakeMediaStore) CurrentProfilePhotosKind(_ context.Context, _ domain.PeerType, _ []int64, _ domain.ProfilePhotoKind) (map[int64]domain.ProfilePhotoRef, error) {
|
||||
return map[int64]domain.ProfilePhotoRef{}, nil
|
||||
}
|
||||
func (f *fakeMediaStore) ListProfilePhotos(_ context.Context, _ domain.PeerType, _ int64, _, _ int, _ int64) ([]int64, int, error) {
|
||||
func (f *fakeMediaStore) ListProfilePhotosKind(_ context.Context, _ domain.PeerType, _ int64, _ domain.ProfilePhotoKind, _, _ int, _ int64) ([]int64, int, error) {
|
||||
return nil, 0, nil
|
||||
}
|
||||
func (f *fakeMediaStore) ListProfilePhotosKind(_ context.Context, _ domain.PeerType, _ int64, _ domain.ProfilePhotoKind, _, _ int, _ int64) ([]int64, int, error) {
|
||||
func (f *fakeMediaStore) ListProfilePhotoDetailsKind(_ context.Context, _ domain.PeerType, _ int64, _ domain.ProfilePhotoKind, _, _ int, _ int64) ([]domain.Photo, int, error) {
|
||||
return nil, 0, nil
|
||||
}
|
||||
func (f *fakeMediaStore) DeleteProfilePhotos(_ context.Context, _ domain.PeerType, _ int64, _ []int64) ([]int64, error) {
|
||||
|
|
@ -252,6 +363,146 @@ func TestSeedMediaRepairsPartialReactionBlobs(t *testing.T) {
|
|||
}
|
||||
}
|
||||
|
||||
func TestSeedCustomEmojiTGSWithoutThumbGetsSyntheticPreview(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
seedDir := t.TempDir()
|
||||
const sourceID int64 = 4444444
|
||||
writeStatusPackWithoutThumbSeed(t, seedDir, sourceID, 17)
|
||||
|
||||
media := newFakeMediaStore()
|
||||
blobs, err := NewLocalFS(t.TempDir())
|
||||
if err != nil {
|
||||
t.Fatalf("local fs: %v", err)
|
||||
}
|
||||
svc := NewService(media, blobs, 2)
|
||||
stats, err := svc.SeedMedia(ctx, seedDir, 0)
|
||||
if err != nil {
|
||||
t.Fatalf("seed media: %v", err)
|
||||
}
|
||||
if stats.StickerSets != 1 || stats.Documents != 1 || stats.Blobs != 2 || stats.Skipped {
|
||||
t.Fatalf("stats = %+v, want one set, one doc, main blob plus synthetic preview", stats)
|
||||
}
|
||||
|
||||
set, ok, err := media.GetStickerSetByShortName(ctx, "StatusPack")
|
||||
if err != nil || !ok {
|
||||
t.Fatalf("StatusPack ok=%v err=%v", ok, err)
|
||||
}
|
||||
doc, ok, err := media.GetDocument(ctx, set.DocumentIDs[0])
|
||||
if err != nil || !ok {
|
||||
t.Fatalf("StatusPack document ok=%v err=%v", ok, err)
|
||||
}
|
||||
if !seedDocumentHasAttribute(doc.Attributes, domain.DocAttrCustomEmoji) {
|
||||
t.Fatalf("document attributes = %+v, want custom emoji", doc.Attributes)
|
||||
}
|
||||
thumb, ok := findCachedThumb(doc.Thumbs)
|
||||
if !ok {
|
||||
t.Fatalf("document thumbs = %+v, want synthetic cached preview", doc.Thumbs)
|
||||
}
|
||||
if thumb.Type != seedSyntheticDocumentThumbType || thumb.W != 1 || thumb.H != 1 || len(thumb.Bytes) == 0 {
|
||||
t.Fatalf("synthetic thumb = %+v, want 1x1 cached %q thumb", thumb, seedSyntheticDocumentThumbType)
|
||||
}
|
||||
blob, ok, err := media.GetFileBlob(ctx, fmt.Sprintf("doc:%d:%s", doc.ID, thumb.Type))
|
||||
if err != nil || !ok {
|
||||
t.Fatalf("synthetic thumb blob ok=%v err=%v", ok, err)
|
||||
}
|
||||
if blob.MimeType != "image/png" {
|
||||
t.Fatalf("synthetic thumb blob mime = %q, want image/png", blob.MimeType)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSeedMediaRepairsCustomEmojiTGSWithoutThumb(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
seedDir := t.TempDir()
|
||||
const sourceID int64 = 5555555
|
||||
const setHash = 23
|
||||
writeStatusPackWithoutThumbSeed(t, seedDir, sourceID, setHash)
|
||||
|
||||
media := newFakeMediaStore()
|
||||
if err := media.PutDocument(ctx, domain.Document{
|
||||
ID: sourceID,
|
||||
MimeType: "application/x-tgsticker",
|
||||
Attributes: []domain.DocumentAttribute{{
|
||||
Kind: domain.DocAttrCustomEmoji,
|
||||
Alt: "\U0001f44b",
|
||||
TextColor: true,
|
||||
}},
|
||||
}); err != nil {
|
||||
t.Fatalf("put stale document: %v", err)
|
||||
}
|
||||
if err := media.PutStickerSet(ctx, domain.StickerSet{
|
||||
ID: 773947703670341676,
|
||||
AccessHash: 1,
|
||||
ShortName: "StatusPack",
|
||||
Title: "Status Pack",
|
||||
Hash: setHash,
|
||||
Kind: domain.StickerSetKindEmoji,
|
||||
Emojis: true,
|
||||
DocumentIDs: []int64{
|
||||
sourceID,
|
||||
},
|
||||
}); err != nil {
|
||||
t.Fatalf("put stale sticker set: %v", err)
|
||||
}
|
||||
|
||||
blobs, err := NewLocalFS(t.TempDir())
|
||||
if err != nil {
|
||||
t.Fatalf("local fs: %v", err)
|
||||
}
|
||||
svc := NewService(media, blobs, 2)
|
||||
stats, err := svc.SeedMedia(ctx, seedDir, 0)
|
||||
if err != nil {
|
||||
t.Fatalf("repair seed: %v", err)
|
||||
}
|
||||
if stats.StickerSets != 1 || stats.Documents != 1 || stats.Blobs != 2 || stats.Skipped {
|
||||
t.Fatalf("repair stats = %+v, want forced reimport", stats)
|
||||
}
|
||||
doc, ok, err := media.GetDocument(ctx, sourceID)
|
||||
if err != nil || !ok {
|
||||
t.Fatalf("repaired document ok=%v err=%v", ok, err)
|
||||
}
|
||||
if _, ok := findCachedThumb(doc.Thumbs); !ok {
|
||||
t.Fatalf("repaired document thumbs = %+v, want synthetic cached preview", doc.Thumbs)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSeedMediaSkipsUnchangedEffectsDocuments(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
seedDir := t.TempDir()
|
||||
const sourceID int64 = 6666666
|
||||
writeEffectsSeed(t, seedDir, sourceID)
|
||||
|
||||
media := newFakeMediaStore()
|
||||
blobs, err := NewLocalFS(t.TempDir())
|
||||
if err != nil {
|
||||
t.Fatalf("local fs: %v", err)
|
||||
}
|
||||
svc := NewService(media, blobs, 2)
|
||||
first, err := svc.SeedMedia(ctx, seedDir, 0)
|
||||
if err != nil {
|
||||
t.Fatalf("first seed: %v", err)
|
||||
}
|
||||
if first.Effects != 1 || first.Documents != 1 || first.Blobs != 1 {
|
||||
t.Fatalf("first stats = %+v, want one imported effect document/blob", first)
|
||||
}
|
||||
|
||||
second, err := svc.SeedMedia(ctx, seedDir, 0)
|
||||
if err != nil {
|
||||
t.Fatalf("second seed: %v", err)
|
||||
}
|
||||
if second.Effects != 1 || second.Documents != 0 || second.Blobs != 0 {
|
||||
t.Fatalf("second stats = %+v, want effects catalog loaded without document/blob import", second)
|
||||
}
|
||||
|
||||
delete(media.blobs, fmt.Sprintf("doc:%d", sourceID))
|
||||
repaired, err := svc.SeedMedia(ctx, seedDir, 0)
|
||||
if err != nil {
|
||||
t.Fatalf("repair seed: %v", err)
|
||||
}
|
||||
if repaired.Effects != 1 || repaired.Documents != 1 || repaired.Blobs != 1 {
|
||||
t.Fatalf("repair stats = %+v, want missing blob to force reimport", repaired)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSeedMediaFromRealExport(t *testing.T) {
|
||||
seedDir := os.Getenv("TELESRV_REAL_STICKER_SEED_DIR")
|
||||
if seedDir == "" {
|
||||
|
|
@ -355,6 +606,37 @@ func TestSeedMediaFromRealExport(t *testing.T) {
|
|||
}
|
||||
}
|
||||
|
||||
func writeStatusPackWithoutThumbSeed(t *testing.T, seedDir string, sourceID int64, setHash int) {
|
||||
t.Helper()
|
||||
setDir := filepath.Join(seedDir, "telegram_emoji_export", "StatusPack_773947703670341676")
|
||||
stickersDir := filepath.Join(setDir, "stickers")
|
||||
if err := os.MkdirAll(stickersDir, 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
raw := fmt.Sprintf(`{"result":{"set":{"id":773947703670341676,"access_hash":1,"title":"Status Pack","short_name":"StatusPack","count":1,"hash":%d,"emojis":true,"packs":[{"emoticon":"👋","documents":[%d]}]},"packs":[{"emoticon":"👋","documents":[%d]}],"documents":[{"id":%d,"access_hash":2,"file_reference":"","date":"2026-06-29T00:00:00Z","mime_type":"application/x-tgsticker","size":4,"dc_id":4,"attributes":[{"_":"DocumentAttributeImageSize","w":512,"h":512},{"_":"DocumentAttributeCustomEmoji","alt":"👋","text_color":true,"stickerset":{"id":773947703670341676,"access_hash":1}},{"_":"DocumentAttributeFilename","file_name":"AnimatedSticker.tgs"}],"thumbs":[]}]}}`, setHash, sourceID, sourceID, sourceID)
|
||||
if err := os.WriteFile(filepath.Join(setDir, "set_info.json"), []byte(raw), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := os.WriteFile(filepath.Join(stickersDir, fmt.Sprintf("status_%d.tgs", sourceID)), []byte("tgs!"), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
func writeEffectsSeed(t *testing.T, seedDir string, sourceID int64) {
|
||||
t.Helper()
|
||||
docsDir := filepath.Join(seedDir, "telegram_effects_export", "documents")
|
||||
if err := os.MkdirAll(docsDir, 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
raw := fmt.Sprintf(`{"result":{"effects":[{"id":77,"emoticon":"🔥","effect_sticker_id":%d}],"documents":[{"id":%d,"access_hash":2,"file_reference":"","date":"2026-06-29T00:00:00Z","mime_type":"application/x-tgsticker","size":4,"dc_id":4,"attributes":[{"_":"DocumentAttributeImageSize","w":512,"h":512},{"_":"DocumentAttributeFilename","file_name":"effect.tgs"}],"thumbs":[]}]}}`, sourceID, sourceID)
|
||||
if err := os.WriteFile(filepath.Join(seedDir, "telegram_effects_export", "effects.json"), []byte(raw), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := os.WriteFile(filepath.Join(docsDir, fmt.Sprintf("effect_%d.tgs", sourceID)), []byte("tgs!"), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSeedDocumentStorageIDNormalizesExternalIDs(t *testing.T) {
|
||||
const sourceID int64 = 5382305375846410902
|
||||
const want int64 = 1382305375846410902
|
||||
|
|
|
|||
|
|
@ -1,11 +1,19 @@
|
|||
package files
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"fmt"
|
||||
"hash"
|
||||
"io"
|
||||
"time"
|
||||
|
||||
"telesrv/internal/domain"
|
||||
"telesrv/internal/store"
|
||||
|
||||
"go.uber.org/zap"
|
||||
"golang.org/x/sync/singleflight"
|
||||
)
|
||||
|
||||
// 上传分片上限:与 Telegram 客户端约定一致(单片 ≤512KB;分片总数有上限防止 OOM)。
|
||||
|
|
@ -14,6 +22,15 @@ const (
|
|||
MaxUploadParts = 8000 // 512KB * 8000 ≈ 4GB 理论上限,足够主路径媒体
|
||||
)
|
||||
|
||||
const (
|
||||
DefaultUploadPartTTL = 24 * time.Hour
|
||||
DefaultUploadPartGCInterval = 30 * time.Minute
|
||||
DefaultUploadPartGCBatch = 10000
|
||||
DefaultUploadInFlightMaxBytes = int64(MaxUploadPartBytes) * int64(MaxUploadParts)
|
||||
DefaultUploadInFlightMaxParts = MaxUploadParts
|
||||
DefaultUploadInFlightMaxFiles = 64
|
||||
)
|
||||
|
||||
// blobMetaCacheCapacity 是 location_key→FileBlob 元数据 LRU 容量(每项约百字节,约 13MB)。
|
||||
const blobMetaCacheCapacity = 1 << 16
|
||||
|
||||
|
|
@ -21,28 +38,104 @@ const blobMetaCacheCapacity = 1 << 16
|
|||
const (
|
||||
blobBytesCacheMaxEntryBytes = 256 << 10 // 256KB
|
||||
blobBytesCacheMaxBytes = 64 << 20 // 64MB
|
||||
// stickerSetNegativeCacheTTL 是未找到贴纸集的负缓存有效期:未 seed 的 short_name 会被客户端
|
||||
// 反复 getStickerSet,这里短时缓存 not-found 短路掉 PG。短 TTL 保证运行时新增集合最多滞后这么久。
|
||||
stickerSetNegativeCacheTTL = 30 * time.Second
|
||||
)
|
||||
|
||||
// Service 实现 upload 分片累积、blob 落盘、getFile 下载,并把上传文件组装成 Photo / Document。
|
||||
type Service struct {
|
||||
media store.MediaStore
|
||||
blobs BlobBackend
|
||||
dc int
|
||||
blobCache *blobMetaCache
|
||||
byteCache *blobBytesCache
|
||||
stickerSetCache *stickerSetFullCache
|
||||
media store.MediaStore
|
||||
blobs BlobBackend
|
||||
uploadParts UploadPartBackend
|
||||
dc int
|
||||
log *zap.Logger
|
||||
thumbs VideoThumbnailer
|
||||
thumbsSet bool
|
||||
blobCache *blobMetaCache
|
||||
byteCache *blobBytesCache
|
||||
// blobMetaSF/blobBytesSF 合并对同一热 blob 的并发首次访问:否则每个并发 getFile 都各打
|
||||
// 一发 PG GetFileBlob + backend GetRange(热门贴纸/reaction/头像被大量用户同时拉时尤甚)。
|
||||
blobMetaSF singleflight.Group
|
||||
blobBytesSF singleflight.Group
|
||||
stickerSetCache *stickerSetFullCache
|
||||
stickerSetNegCache *stickerSetNegativeCache
|
||||
uploadQuota domain.UploadPartQuota
|
||||
mapTiles *mapTileProxy
|
||||
externalMedia *externalMediaFetcher
|
||||
webpage *webpageFetcher
|
||||
// effects 是消息发送特效目录(messages.getAvailableEffects)。全局静态,启动 seedEffects
|
||||
// 一次写入后只读,故无锁——与各 read-model 缓存一样在服务就绪前完成填充。
|
||||
// effectsHash 在 seed 时算一次,handler 直接比对返回 NotModified,无需每次 RPC 重算。
|
||||
effects []domain.AvailableEffect
|
||||
effectsHash int
|
||||
}
|
||||
|
||||
// Option 配置 files 服务的可选能力。
|
||||
type Option func(*Service)
|
||||
|
||||
// WithLogger 注入日志器。未注入时使用 no-op logger。
|
||||
func WithLogger(log *zap.Logger) Option {
|
||||
return func(s *Service) {
|
||||
if log != nil {
|
||||
s.log = log
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// WithVideoThumbnailer 覆盖视频缩略图生成器。传 nil 可显式关闭服务端抽帧 fallback。
|
||||
func WithVideoThumbnailer(thumbnailer VideoThumbnailer) Option {
|
||||
return func(s *Service) {
|
||||
s.thumbs = thumbnailer
|
||||
s.thumbsSet = true
|
||||
}
|
||||
}
|
||||
|
||||
// WithUploadPartQuota 覆盖用户级 in-flight 上传分片配额;字段 <=0 表示该维度不限制。
|
||||
func WithUploadPartQuota(quota domain.UploadPartQuota) Option {
|
||||
return func(s *Service) {
|
||||
s.uploadQuota = quota
|
||||
}
|
||||
}
|
||||
|
||||
// NewService 创建 files 服务。dc 是本 server 的 DC id,写入新建 document/photo 的 dc_id。
|
||||
func NewService(media store.MediaStore, blobs BlobBackend, dc int) *Service {
|
||||
return &Service{
|
||||
media: media,
|
||||
blobs: blobs,
|
||||
dc: dc,
|
||||
blobCache: newBlobMetaCache(blobMetaCacheCapacity),
|
||||
byteCache: newBlobBytesCache(blobBytesCacheMaxBytes),
|
||||
stickerSetCache: newStickerSetFullCache(),
|
||||
func NewService(media store.MediaStore, blobs BlobBackend, dc int, opts ...Option) *Service {
|
||||
s := &Service{
|
||||
media: media,
|
||||
blobs: blobs,
|
||||
dc: dc,
|
||||
log: zap.NewNop(),
|
||||
blobCache: newBlobMetaCache(blobMetaCacheCapacity),
|
||||
byteCache: newBlobBytesCache(blobBytesCacheMaxBytes),
|
||||
stickerSetCache: newStickerSetFullCache(),
|
||||
stickerSetNegCache: newStickerSetNegativeCache(stickerSetNegativeCacheTTL),
|
||||
uploadQuota: domain.UploadPartQuota{
|
||||
MaxBytes: DefaultUploadInFlightMaxBytes,
|
||||
MaxParts: DefaultUploadInFlightMaxParts,
|
||||
MaxFiles: DefaultUploadInFlightMaxFiles,
|
||||
},
|
||||
}
|
||||
if partBackend, ok := blobs.(UploadPartBackend); ok {
|
||||
s.uploadParts = partBackend
|
||||
}
|
||||
for _, opt := range opts {
|
||||
if opt != nil {
|
||||
opt(s)
|
||||
}
|
||||
}
|
||||
if s.mapTiles != nil {
|
||||
// 选项应用顺序无关:logger 在全部 Option 跑完后统一注入。
|
||||
s.mapTiles.log = s.log
|
||||
}
|
||||
if !s.thumbsSet {
|
||||
thumbnailer, err := NewFFmpegVideoThumbnailer()
|
||||
if err != nil {
|
||||
s.log.Warn("ffmpeg not found; server-side video thumbnail fallback disabled", zap.Error(err))
|
||||
} else {
|
||||
s.thumbs = thumbnailer
|
||||
}
|
||||
}
|
||||
return s
|
||||
}
|
||||
|
||||
// SaveFilePart 累积一个 small file 分片。
|
||||
|
|
@ -50,12 +143,12 @@ func (s *Service) SaveFilePart(ctx context.Context, ownerUserID, fileID int64, p
|
|||
if err := validatePart(part, len(bytes)); err != nil {
|
||||
return false, err
|
||||
}
|
||||
if err := s.media.SaveFilePart(ctx, domain.UploadPart{
|
||||
if err := s.saveFilePart(ctx, domain.UploadPart{
|
||||
OwnerUserID: ownerUserID,
|
||||
FileID: fileID,
|
||||
Part: part,
|
||||
Bytes: bytes,
|
||||
}); err != nil {
|
||||
Size: int64(len(bytes)),
|
||||
}, bytes); err != nil {
|
||||
return false, err
|
||||
}
|
||||
return true, nil
|
||||
|
|
@ -69,37 +162,142 @@ func (s *Service) SaveBigFilePart(ctx context.Context, ownerUserID, fileID int64
|
|||
if totalParts <= 0 || totalParts > MaxUploadParts {
|
||||
return false, domain.ErrFilePartsInvalid
|
||||
}
|
||||
if err := s.media.SaveFilePart(ctx, domain.UploadPart{
|
||||
if err := s.saveFilePart(ctx, domain.UploadPart{
|
||||
OwnerUserID: ownerUserID,
|
||||
FileID: fileID,
|
||||
Part: part,
|
||||
TotalParts: totalParts,
|
||||
Big: true,
|
||||
Bytes: bytes,
|
||||
}); err != nil {
|
||||
Size: int64(len(bytes)),
|
||||
}, bytes); err != nil {
|
||||
return false, err
|
||||
}
|
||||
return true, nil
|
||||
}
|
||||
|
||||
func (s *Service) saveFilePart(ctx context.Context, part domain.UploadPart, bytes []byte) error {
|
||||
if s.uploadParts == nil {
|
||||
return fmt.Errorf("upload part backend not configured")
|
||||
}
|
||||
slot, err := s.checkUploadPartQuota(ctx, part)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
obj, err := s.uploadParts.PutUploadPart(ctx, part.OwnerUserID, part.FileID, part.Part, bytes)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
part.Backend = obj.Backend
|
||||
part.ObjectKey = obj.ObjectKey
|
||||
part.Size = obj.Size
|
||||
part.SHA256 = obj.SHA256
|
||||
if err := s.media.SaveFilePart(ctx, part); err != nil {
|
||||
_ = s.uploadParts.DeleteUploadPart(ctx, obj.ObjectKey)
|
||||
return err
|
||||
}
|
||||
if slot.Found && slot.ObjectKey != "" && slot.ObjectKey != obj.ObjectKey {
|
||||
if err := s.uploadParts.DeleteUploadPart(ctx, slot.ObjectKey); err != nil {
|
||||
s.log.Warn("delete replaced upload part failed", zap.String("object_key", slot.ObjectKey), zap.Error(err))
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *Service) checkUploadPartQuota(ctx context.Context, part domain.UploadPart) (domain.UploadPartSlot, error) {
|
||||
slot, err := s.media.UploadPartSlot(ctx, part.OwnerUserID, part.FileID, part.Part)
|
||||
if err != nil {
|
||||
return domain.UploadPartSlot{}, err
|
||||
}
|
||||
quota := s.uploadQuota
|
||||
if quota.MaxBytes <= 0 && quota.MaxParts <= 0 && quota.MaxFiles <= 0 {
|
||||
return slot, nil
|
||||
}
|
||||
usage, err := s.media.UploadPartUsage(ctx, part.OwnerUserID)
|
||||
if err != nil {
|
||||
return domain.UploadPartSlot{}, err
|
||||
}
|
||||
next := usage
|
||||
next.Bytes += part.Size - slot.ExistingBytes
|
||||
if !slot.Found {
|
||||
next.Parts++
|
||||
}
|
||||
if slot.FileParts == 0 {
|
||||
next.Files++
|
||||
}
|
||||
if quota.MaxBytes > 0 && next.Bytes > quota.MaxBytes {
|
||||
return domain.UploadPartSlot{}, domain.ErrUploadQuotaExceeded
|
||||
}
|
||||
if quota.MaxParts > 0 && next.Parts > quota.MaxParts {
|
||||
return domain.UploadPartSlot{}, domain.ErrUploadQuotaExceeded
|
||||
}
|
||||
if quota.MaxFiles > 0 && next.Files > quota.MaxFiles {
|
||||
return domain.UploadPartSlot{}, domain.ErrUploadQuotaExceeded
|
||||
}
|
||||
return slot, nil
|
||||
}
|
||||
|
||||
// DeleteExpiredUploadParts 清理超过保留期仍未组装的 transient 上传分片。
|
||||
func (s *Service) DeleteExpiredUploadParts(ctx context.Context, before time.Time, limit int) (int64, error) {
|
||||
if limit <= 0 {
|
||||
return 0, nil
|
||||
}
|
||||
keys, err := s.media.DeleteExpiredUploadParts(ctx, before, limit)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
if err := s.deleteUploadPartObjects(ctx, keys); err != nil {
|
||||
return int64(len(keys)), err
|
||||
}
|
||||
var orphanDeleted int64
|
||||
if s.uploadParts != nil {
|
||||
n, err := s.uploadParts.DeleteExpiredUploadParts(ctx, before, limit)
|
||||
if err != nil {
|
||||
return int64(len(keys)), err
|
||||
}
|
||||
orphanDeleted = n
|
||||
}
|
||||
return int64(len(keys)) + orphanDeleted, nil
|
||||
}
|
||||
|
||||
// GetFile 按 location_key 取一段 blob 内容。found=false 表示该 location 无对应 blob。
|
||||
// 元数据走进程内 LRU(消除每 chunk 一次 PG 查);小 blob 全量字节进 LRU,供 sticker /
|
||||
// reaction / thumbnail 热路径直接内存切片;大 blob 仍按 offset/limit 段读。
|
||||
type blobMetaResult struct {
|
||||
blob domain.FileBlob
|
||||
found bool
|
||||
}
|
||||
|
||||
type blobBytesResult struct {
|
||||
data []byte
|
||||
total int64
|
||||
cacheable bool
|
||||
}
|
||||
|
||||
func (s *Service) GetFile(ctx context.Context, req domain.FileDownloadRequest) (domain.FileChunk, bool, error) {
|
||||
blob, ok := s.blobCache.get(req.LocationKey)
|
||||
if !ok {
|
||||
var (
|
||||
found bool
|
||||
err error
|
||||
)
|
||||
blob, found, err = s.media.GetFileBlob(ctx, req.LocationKey)
|
||||
// 同一 location_key 的并发首访合并成一次 PG GetFileBlob。
|
||||
v, err, _ := s.blobMetaSF.Do(req.LocationKey, func() (any, error) {
|
||||
if cached, ok := s.blobCache.get(req.LocationKey); ok {
|
||||
return blobMetaResult{blob: cached, found: true}, nil
|
||||
}
|
||||
b, found, err := s.media.GetFileBlob(ctx, req.LocationKey)
|
||||
if err != nil {
|
||||
return blobMetaResult{}, err
|
||||
}
|
||||
if found {
|
||||
s.blobCache.put(req.LocationKey, b)
|
||||
}
|
||||
return blobMetaResult{blob: b, found: found}, nil
|
||||
})
|
||||
if err != nil {
|
||||
return domain.FileChunk{}, false, err
|
||||
}
|
||||
if !found {
|
||||
res := v.(blobMetaResult)
|
||||
if !res.found {
|
||||
return domain.FileChunk{}, false, nil
|
||||
}
|
||||
s.blobCache.put(req.LocationKey, blob)
|
||||
blob = res.blob
|
||||
}
|
||||
if blob.Size > 0 && blob.Size <= blobBytesCacheMaxEntryBytes {
|
||||
if data, ok := s.byteCache.get(blob.ObjectKey); ok {
|
||||
|
|
@ -109,18 +307,33 @@ func (s *Service) GetFile(ctx context.Context, req domain.FileDownloadRequest) (
|
|||
Total: int64(len(data)),
|
||||
}, true, nil
|
||||
}
|
||||
data, total, err := s.blobs.GetRange(ctx, blob.ObjectKey, 0, blobBytesCacheMaxEntryBytes+1)
|
||||
// 同一 object_key 的小 blob 并发首访合并成一次 backend 全量读 + 一次 byteCache 填充。
|
||||
v, err, _ := s.blobBytesSF.Do(blob.ObjectKey, func() (any, error) {
|
||||
if cached, ok := s.byteCache.get(blob.ObjectKey); ok {
|
||||
return blobBytesResult{data: cached, total: int64(len(cached)), cacheable: true}, nil
|
||||
}
|
||||
data, total, err := s.blobs.GetRange(ctx, blob.ObjectKey, 0, blobBytesCacheMaxEntryBytes+1)
|
||||
if err != nil {
|
||||
return blobBytesResult{}, err
|
||||
}
|
||||
if total <= blobBytesCacheMaxEntryBytes && int64(len(data)) == total {
|
||||
s.byteCache.put(blob.ObjectKey, data)
|
||||
return blobBytesResult{data: data, total: total, cacheable: true}, nil
|
||||
}
|
||||
return blobBytesResult{cacheable: false}, nil
|
||||
})
|
||||
if err != nil {
|
||||
return domain.FileChunk{}, false, fmt.Errorf("read blob %q: %w", blob.LocationKey, err)
|
||||
}
|
||||
if total <= blobBytesCacheMaxEntryBytes && int64(len(data)) == total {
|
||||
s.byteCache.put(blob.ObjectKey, data)
|
||||
// res.data 在并发 caller 间只读共享,sliceBlobBytes 各自拷贝出自己的分片,安全。
|
||||
if res := v.(blobBytesResult); res.cacheable {
|
||||
return domain.FileChunk{
|
||||
Bytes: sliceBlobBytes(data, req.Offset, int64(req.Limit)),
|
||||
Bytes: sliceBlobBytes(res.data, req.Offset, int64(req.Limit)),
|
||||
MimeType: blob.MimeType,
|
||||
Total: total,
|
||||
Total: res.total,
|
||||
}, true, nil
|
||||
}
|
||||
// 大小不符/超限:落到下面的按需 range 读(与原行为一致)。
|
||||
}
|
||||
data, total, err := s.blobs.GetRange(ctx, blob.ObjectKey, req.Offset, int64(req.Limit))
|
||||
if err != nil {
|
||||
|
|
@ -170,6 +383,11 @@ func (s *Service) ResolveStickerSet(ctx context.Context, ref domain.StickerSetRe
|
|||
if set, docs, ok := s.stickerSetCache.get(ref); ok {
|
||||
return set, docs, true, nil
|
||||
}
|
||||
// 负缓存:已 seed 集启动即进正缓存(WarmCaches),能走到这里的 miss 多是「未 seed 的 short_name」
|
||||
// 被客户端反复请求。TTL 内直接当 not-found 短路,避免每次都打 PG GetStickerSetByShortName。
|
||||
if s.stickerSetNegCache != nil && s.stickerSetNegCache.has(ref) {
|
||||
return domain.StickerSet{}, nil, false, nil
|
||||
}
|
||||
var (
|
||||
set domain.StickerSet
|
||||
found bool
|
||||
|
|
@ -186,6 +404,9 @@ func (s *Service) ResolveStickerSet(ctx context.Context, ref domain.StickerSetRe
|
|||
return domain.StickerSet{}, nil, false, nil
|
||||
}
|
||||
if err != nil || !found {
|
||||
if err == nil && !found && s.stickerSetNegCache != nil {
|
||||
s.stickerSetNegCache.put(ref)
|
||||
}
|
||||
return domain.StickerSet{}, nil, found, err
|
||||
}
|
||||
docs, err := s.media.GetDocuments(ctx, set.DocumentIDs)
|
||||
|
|
@ -215,33 +436,226 @@ func orderDocuments(docs []domain.Document, ids []int64) []domain.Document {
|
|||
// assembleUpload 把已上传分片按 part 顺序拼成完整字节,并清理分片。
|
||||
// expectedParts>0 时校验分片连续且齐全。
|
||||
func (s *Service) assembleUpload(ctx context.Context, ownerUserID, fileID int64, expectedParts int) ([]byte, error) {
|
||||
parts, err := s.media.LoadFileParts(ctx, ownerUserID, fileID)
|
||||
parts, _, err := s.loadAndValidateUploadParts(ctx, ownerUserID, fileID, expectedParts)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if len(parts) == 0 {
|
||||
return nil, domain.ErrFilePartsInvalid
|
||||
}
|
||||
if expectedParts > 0 && len(parts) != expectedParts {
|
||||
return nil, domain.ErrFilePartsInvalid
|
||||
}
|
||||
total := 0
|
||||
for i, p := range parts {
|
||||
if p.Part != i {
|
||||
return nil, domain.ErrFilePartsInvalid // 缺片或乱序
|
||||
}
|
||||
total += len(p.Bytes)
|
||||
}
|
||||
buf := make([]byte, 0, total)
|
||||
buf := make([]byte, 0, uploadPartsTotalSize(parts))
|
||||
for _, p := range parts {
|
||||
buf = append(buf, p.Bytes...)
|
||||
if s.uploadParts == nil {
|
||||
return nil, fmt.Errorf("upload part backend not configured")
|
||||
}
|
||||
data, err := s.uploadParts.GetUploadPart(ctx, p.ObjectKey)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("read upload part %d: %w", p.Part, err)
|
||||
}
|
||||
if err := validateUploadPartBytes(p, data); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
buf = append(buf, data...)
|
||||
}
|
||||
if err := s.media.DeleteFileParts(ctx, ownerUserID, fileID); err != nil {
|
||||
if err := s.cleanupUploadParts(ctx, ownerUserID, fileID); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return buf, nil
|
||||
}
|
||||
|
||||
type assembledUploadBlob struct {
|
||||
ObjectKey string
|
||||
Size int64
|
||||
SHA256 []byte
|
||||
}
|
||||
|
||||
// assembleUploadBlob 把上传分片流式写入正式 blob。调用方应在 durable media 元数据
|
||||
// 成功提交后调用 cleanupUploadParts,避免 metadata 写失败时丢失可重试的上传分片。
|
||||
func (s *Service) assembleUploadBlob(ctx context.Context, ownerUserID, fileID int64, expectedParts int) (assembledUploadBlob, error) {
|
||||
parts, _, err := s.loadAndValidateUploadParts(ctx, ownerUserID, fileID, expectedParts)
|
||||
if err != nil {
|
||||
return assembledUploadBlob{}, err
|
||||
}
|
||||
if s.uploadParts == nil {
|
||||
return assembledUploadBlob{}, fmt.Errorf("upload part backend not configured")
|
||||
}
|
||||
reader := &uploadPartsReader{
|
||||
ctx: ctx,
|
||||
backend: s.uploadParts,
|
||||
parts: parts,
|
||||
}
|
||||
defer reader.Close()
|
||||
objectKey, size, sum, err := s.blobs.PutReader(ctx, reader)
|
||||
if err != nil {
|
||||
return assembledUploadBlob{}, err
|
||||
}
|
||||
return assembledUploadBlob{
|
||||
ObjectKey: objectKey,
|
||||
Size: size,
|
||||
SHA256: sum,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (s *Service) loadAndValidateUploadParts(ctx context.Context, ownerUserID, fileID int64, expectedParts int) ([]domain.UploadPart, int64, error) {
|
||||
parts, err := s.media.LoadFileParts(ctx, ownerUserID, fileID)
|
||||
if err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
if len(parts) == 0 {
|
||||
return nil, 0, domain.ErrFilePartsInvalid
|
||||
}
|
||||
if expectedParts > 0 && len(parts) != expectedParts {
|
||||
return nil, 0, domain.ErrFilePartsInvalid
|
||||
}
|
||||
var total int64
|
||||
for i, p := range parts {
|
||||
if p.Part != i {
|
||||
return nil, 0, domain.ErrFilePartsInvalid // 缺片或乱序
|
||||
}
|
||||
if p.Size <= 0 || p.Size > MaxUploadPartBytes || p.ObjectKey == "" {
|
||||
return nil, 0, domain.ErrFilePartsInvalid
|
||||
}
|
||||
total += p.Size
|
||||
if total > DefaultUploadInFlightMaxBytes {
|
||||
return nil, 0, domain.ErrFilePartsInvalid
|
||||
}
|
||||
}
|
||||
return parts, total, nil
|
||||
}
|
||||
|
||||
func uploadPartsTotalSize(parts []domain.UploadPart) int {
|
||||
var total int64
|
||||
for _, p := range parts {
|
||||
total += p.Size
|
||||
}
|
||||
return int(total)
|
||||
}
|
||||
|
||||
func validateUploadPartBytes(part domain.UploadPart, data []byte) error {
|
||||
if int64(len(data)) != part.Size {
|
||||
return domain.ErrFilePartsInvalid
|
||||
}
|
||||
if len(part.SHA256) > 0 {
|
||||
sum := sha256.Sum256(data)
|
||||
if !bytes.Equal(sum[:], part.SHA256) {
|
||||
return domain.ErrFilePartsInvalid
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *Service) cleanupUploadParts(ctx context.Context, ownerUserID, fileID int64) error {
|
||||
keys, err := s.media.DeleteFileParts(ctx, ownerUserID, fileID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := s.deleteUploadPartObjects(ctx, keys); err != nil {
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
type uploadPartsReader struct {
|
||||
ctx context.Context
|
||||
backend UploadPartBackend
|
||||
parts []domain.UploadPart
|
||||
index int
|
||||
current io.ReadCloser
|
||||
currentRead int64
|
||||
currentHash hash.Hash
|
||||
}
|
||||
|
||||
func (r *uploadPartsReader) Read(buf []byte) (int, error) {
|
||||
for {
|
||||
if r.current == nil {
|
||||
if r.index >= len(r.parts) {
|
||||
return 0, io.EOF
|
||||
}
|
||||
select {
|
||||
case <-r.ctx.Done():
|
||||
return 0, r.ctx.Err()
|
||||
default:
|
||||
}
|
||||
part := r.parts[r.index]
|
||||
rc, err := r.backend.OpenUploadPart(r.ctx, part.ObjectKey)
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("open upload part %d: %w", part.Part, err)
|
||||
}
|
||||
r.current = rc
|
||||
r.currentRead = 0
|
||||
if len(part.SHA256) > 0 {
|
||||
r.currentHash = sha256.New()
|
||||
} else {
|
||||
r.currentHash = nil
|
||||
}
|
||||
}
|
||||
n, err := r.current.Read(buf)
|
||||
if n > 0 {
|
||||
r.currentRead += int64(n)
|
||||
if r.currentHash != nil {
|
||||
_, _ = r.currentHash.Write(buf[:n])
|
||||
}
|
||||
return n, nil
|
||||
}
|
||||
if err == io.EOF {
|
||||
if err := r.finishCurrentPart(); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
continue
|
||||
}
|
||||
if err != nil {
|
||||
_ = r.current.Close()
|
||||
part := r.parts[r.index]
|
||||
r.current = nil
|
||||
return 0, fmt.Errorf("read upload part %d: %w", part.Part, err)
|
||||
}
|
||||
return 0, nil
|
||||
}
|
||||
}
|
||||
|
||||
func (r *uploadPartsReader) finishCurrentPart() error {
|
||||
part := r.parts[r.index]
|
||||
closeErr := r.current.Close()
|
||||
r.current = nil
|
||||
if closeErr != nil {
|
||||
return fmt.Errorf("close upload part %d: %w", part.Part, closeErr)
|
||||
}
|
||||
if r.currentRead != part.Size {
|
||||
return domain.ErrFilePartsInvalid
|
||||
}
|
||||
if r.currentHash != nil && !bytes.Equal(r.currentHash.Sum(nil), part.SHA256) {
|
||||
return domain.ErrFilePartsInvalid
|
||||
}
|
||||
r.currentHash = nil
|
||||
r.currentRead = 0
|
||||
r.index++
|
||||
return nil
|
||||
}
|
||||
|
||||
func (r *uploadPartsReader) Close() error {
|
||||
if r.current == nil {
|
||||
return nil
|
||||
}
|
||||
err := r.current.Close()
|
||||
r.current = nil
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *Service) deleteUploadPartObjects(ctx context.Context, keys []string) error {
|
||||
if len(keys) == 0 {
|
||||
return nil
|
||||
}
|
||||
if s.uploadParts == nil {
|
||||
return fmt.Errorf("upload part backend not configured")
|
||||
}
|
||||
for _, key := range keys {
|
||||
if key == "" {
|
||||
continue
|
||||
}
|
||||
if err := s.uploadParts.DeleteUploadPart(ctx, key); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func validatePart(part, size int) error {
|
||||
if part < 0 || part >= MaxUploadParts {
|
||||
return domain.ErrFilePartInvalid
|
||||
|
|
|
|||
97
internal/app/files/star_gifts_catalog.go
Normal file
97
internal/app/files/star_gifts_catalog.go
Normal file
|
|
@ -0,0 +1,97 @@
|
|||
package files
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
|
||||
"telesrv/internal/domain"
|
||||
)
|
||||
|
||||
// Star gift 目录:从已 seed 的 animated_emoji 集按 emoticon 精选贴纸文档合成(复用文档行与
|
||||
// blob,不复制字节),镜像 EnsureDefaultEmojiStatusSet。目录是静态的,不入库。
|
||||
// 礼物 ID 取明显隔离的常量段避免撞键。
|
||||
|
||||
const starGiftIDBase int64 = 8_888_000_000_000_000
|
||||
|
||||
type starGiftSeed struct {
|
||||
id int64
|
||||
emoticon string
|
||||
stars int64
|
||||
title string
|
||||
}
|
||||
|
||||
// starGiftSeeds 是固定礼物目录(emoticon 需在 animated_emoji 集里,否则该礼物被跳过)。
|
||||
// convert_stars = stars(v1 全额转换,视作用新购 Stars 买入)。
|
||||
var starGiftSeeds = []starGiftSeed{
|
||||
{starGiftIDBase + 1, "❤", 15, "Heart"},
|
||||
{starGiftIDBase + 2, "\U0001f382", 50, "Cake"}, // 🎂
|
||||
{starGiftIDBase + 3, "\U0001f389", 100, "Party"}, // 🎉
|
||||
{starGiftIDBase + 4, "\U0001f525", 250, "Fire"}, // 🔥
|
||||
{starGiftIDBase + 5, "\U0001f3c6", 500, "Trophy"}, // 🏆
|
||||
{starGiftIDBase + 6, "\U0001f48e", 1000, "Diamond"}, // 💎
|
||||
{starGiftIDBase + 7, "\U0001f680", 2500, "Rocket"}, // 🚀
|
||||
}
|
||||
|
||||
// BuildStarGiftCatalog 合成可购买礼物目录:解析每个 seed emoticon 的贴纸文档,跳过未 seed 的。
|
||||
// animated_emoji 未 seed 时返回空目录(客户端显示空礼物面板,购买流仍可对已知 gift_id 工作)。
|
||||
func (s *Service) BuildStarGiftCatalog(ctx context.Context) ([]domain.StarGift, error) {
|
||||
source, found, err := s.media.GetStickerSetBySystemKey(ctx, "animated_emoji")
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("lookup animated_emoji set for star gifts: %w", err)
|
||||
}
|
||||
if !found || len(source.Packs) == 0 {
|
||||
return nil, nil
|
||||
}
|
||||
byEmoticon := make(map[string]int64, len(source.Packs))
|
||||
for _, pack := range source.Packs {
|
||||
key := normalizeStatusEmoticon(pack.Emoticon)
|
||||
if key == "" || len(pack.DocumentIDs) == 0 {
|
||||
continue
|
||||
}
|
||||
if _, ok := byEmoticon[key]; !ok {
|
||||
byEmoticon[key] = pack.DocumentIDs[0]
|
||||
}
|
||||
}
|
||||
// 收集要加载的文档 id(去重)。
|
||||
docIDs := make([]int64, 0, len(starGiftSeeds))
|
||||
chosen := make([]starGiftSeed, 0, len(starGiftSeeds))
|
||||
seen := make(map[int64]struct{})
|
||||
for _, seed := range starGiftSeeds {
|
||||
id, ok := byEmoticon[normalizeStatusEmoticon(seed.emoticon)]
|
||||
if !ok || id == 0 {
|
||||
continue
|
||||
}
|
||||
chosen = append(chosen, seed)
|
||||
if _, dup := seen[id]; !dup {
|
||||
seen[id] = struct{}{}
|
||||
docIDs = append(docIDs, id)
|
||||
}
|
||||
}
|
||||
if len(chosen) == 0 {
|
||||
return nil, nil
|
||||
}
|
||||
docs, err := s.media.GetDocuments(ctx, docIDs)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("load star gift sticker documents: %w", err)
|
||||
}
|
||||
docByID := make(map[int64]domain.Document, len(docs))
|
||||
for _, d := range docs {
|
||||
docByID[d.ID] = d
|
||||
}
|
||||
catalog := make([]domain.StarGift, 0, len(chosen))
|
||||
for _, seed := range chosen {
|
||||
id := byEmoticon[normalizeStatusEmoticon(seed.emoticon)]
|
||||
doc, ok := docByID[id]
|
||||
if !ok || doc.ID == 0 {
|
||||
continue
|
||||
}
|
||||
catalog = append(catalog, domain.StarGift{
|
||||
ID: seed.id,
|
||||
Stars: seed.stars,
|
||||
ConvertStars: seed.stars,
|
||||
Title: seed.title,
|
||||
Sticker: doc,
|
||||
})
|
||||
}
|
||||
return catalog, nil
|
||||
}
|
||||
67
internal/app/files/upload_gc.go
Normal file
67
internal/app/files/upload_gc.go
Normal file
|
|
@ -0,0 +1,67 @@
|
|||
package files
|
||||
|
||||
import (
|
||||
"context"
|
||||
"time"
|
||||
|
||||
"go.uber.org/zap"
|
||||
)
|
||||
|
||||
// UploadPartGCWorker 周期性清理未组装的过期上传分片。
|
||||
type UploadPartGCWorker struct {
|
||||
files *Service
|
||||
logger *zap.Logger
|
||||
ttl time.Duration
|
||||
interval time.Duration
|
||||
batch int
|
||||
}
|
||||
|
||||
func NewUploadPartGCWorker(files *Service, logger *zap.Logger, ttl, interval time.Duration, batch int) *UploadPartGCWorker {
|
||||
if logger == nil {
|
||||
logger = zap.NewNop()
|
||||
}
|
||||
if ttl <= 0 {
|
||||
ttl = DefaultUploadPartTTL
|
||||
}
|
||||
if interval <= 0 {
|
||||
interval = DefaultUploadPartGCInterval
|
||||
}
|
||||
if batch <= 0 {
|
||||
batch = DefaultUploadPartGCBatch
|
||||
}
|
||||
return &UploadPartGCWorker{
|
||||
files: files,
|
||||
logger: logger,
|
||||
ttl: ttl,
|
||||
interval: interval,
|
||||
batch: batch,
|
||||
}
|
||||
}
|
||||
|
||||
func (w *UploadPartGCWorker) Run(ctx context.Context) {
|
||||
w.runOnce(ctx)
|
||||
ticker := time.NewTicker(w.interval)
|
||||
defer ticker.Stop()
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case <-ticker.C:
|
||||
w.runOnce(ctx)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (w *UploadPartGCWorker) runOnce(ctx context.Context) {
|
||||
if w.files == nil {
|
||||
return
|
||||
}
|
||||
deleted, err := w.files.DeleteExpiredUploadParts(ctx, time.Now().Add(-w.ttl), w.batch)
|
||||
if err != nil {
|
||||
w.logger.Warn("清理过期 upload_parts 失败", zap.Error(err))
|
||||
return
|
||||
}
|
||||
if deleted > 0 {
|
||||
w.logger.Info("清理过期 upload_parts 完成", zap.Int64("deleted", deleted))
|
||||
}
|
||||
}
|
||||
149
internal/app/files/upload_parts_test.go
Normal file
149
internal/app/files/upload_parts_test.go
Normal file
|
|
@ -0,0 +1,149 @@
|
|||
package files
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"strconv"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"telesrv/internal/domain"
|
||||
)
|
||||
|
||||
func TestSaveFilePartQuotaTreatsRetryAsOverwrite(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
media := newFakeMediaStore()
|
||||
svc, blobs := newUploadPartTestService(t, media, domain.UploadPartQuota{MaxBytes: 4, MaxParts: 1, MaxFiles: 1})
|
||||
|
||||
if _, err := svc.SaveFilePart(ctx, 10, 100, 0, []byte("1234")); err != nil {
|
||||
t.Fatalf("save first part: %v", err)
|
||||
}
|
||||
firstParts, err := media.LoadFileParts(ctx, 10, 100)
|
||||
if err != nil || len(firstParts) != 1 || firstParts[0].ObjectKey == "" {
|
||||
t.Fatalf("load first part metadata: parts=%+v err=%v", firstParts, err)
|
||||
}
|
||||
firstKey := firstParts[0].ObjectKey
|
||||
if _, err := svc.SaveFilePart(ctx, 10, 100, 0, []byte("1234")); err != nil {
|
||||
t.Fatalf("retry same part should overwrite without extra quota: %v", err)
|
||||
}
|
||||
parts, err := media.LoadFileParts(ctx, 10, 100)
|
||||
if err != nil {
|
||||
t.Fatalf("load parts: %v", err)
|
||||
}
|
||||
if len(parts) != 1 || parts[0].Size != 4 || parts[0].ObjectKey == "" || parts[0].ObjectKey == firstKey {
|
||||
t.Fatalf("parts after retry = %+v", parts)
|
||||
}
|
||||
if _, err := blobs.GetUploadPart(ctx, firstKey); err == nil {
|
||||
t.Fatalf("replaced upload part object %q still exists", firstKey)
|
||||
}
|
||||
data, err := svc.assembleUpload(ctx, 10, 100, 1)
|
||||
if err != nil {
|
||||
t.Fatalf("assemble upload: %v", err)
|
||||
}
|
||||
if string(data) != "1234" {
|
||||
t.Fatalf("assembled data = %q", data)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSaveFilePartQuotaRejectsNewFileOverLimit(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
media := newFakeMediaStore()
|
||||
svc, _ := newUploadPartTestService(t, media, domain.UploadPartQuota{MaxBytes: 8, MaxParts: 4, MaxFiles: 1})
|
||||
|
||||
if _, err := svc.SaveFilePart(ctx, 10, 100, 0, []byte("1234")); err != nil {
|
||||
t.Fatalf("save first file part: %v", err)
|
||||
}
|
||||
_, err := svc.SaveFilePart(ctx, 10, 101, 0, []byte("12"))
|
||||
if !errors.Is(err, domain.ErrUploadQuotaExceeded) {
|
||||
t.Fatalf("save second file err = %v, want ErrUploadQuotaExceeded", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSaveFilePartQuotaRejectsPartAndByteOverLimit(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
media := newFakeMediaStore()
|
||||
svc, _ := newUploadPartTestService(t, media, domain.UploadPartQuota{MaxBytes: 5, MaxParts: 1, MaxFiles: 2})
|
||||
|
||||
if _, err := svc.SaveFilePart(ctx, 10, 100, 0, []byte("1234")); err != nil {
|
||||
t.Fatalf("save first part: %v", err)
|
||||
}
|
||||
_, err := svc.SaveFilePart(ctx, 10, 100, 1, []byte("12"))
|
||||
if !errors.Is(err, domain.ErrUploadQuotaExceeded) {
|
||||
t.Fatalf("save second part err = %v, want ErrUploadQuotaExceeded", err)
|
||||
}
|
||||
_, err = svc.SaveFilePart(ctx, 10, 100, 0, []byte("123456"))
|
||||
if !errors.Is(err, domain.ErrUploadQuotaExceeded) {
|
||||
t.Fatalf("grow retried part err = %v, want ErrUploadQuotaExceeded", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCreateDocumentFromUploadStreamsBodyAndCleansParts(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
media := newFakeMediaStore()
|
||||
local, err := NewLocalFS(t.TempDir())
|
||||
if err != nil {
|
||||
t.Fatalf("NewLocalFS: %v", err)
|
||||
}
|
||||
blobs := &countingUploadPartBackend{LocalFS: local}
|
||||
svc := NewService(media, blobs, 2, WithVideoThumbnailer(nil))
|
||||
|
||||
parts := []string{
|
||||
strings.Repeat("a", 1024),
|
||||
strings.Repeat("b", 1024),
|
||||
strings.Repeat("c", 1024),
|
||||
}
|
||||
for i, part := range parts {
|
||||
if _, err := svc.SaveBigFilePart(ctx, 10, 200, i, len(parts), []byte(part)); err != nil {
|
||||
t.Fatalf("SaveBigFilePart %d: %v", i, err)
|
||||
}
|
||||
}
|
||||
doc, err := svc.CreateDocumentFromUpload(ctx,
|
||||
domain.UploadedFileRef{OwnerUserID: 10, FileID: 200, Parts: len(parts), Name: "large.bin", Big: true},
|
||||
domain.DocumentSpec{MimeType: "application/octet-stream"},
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("CreateDocumentFromUpload: %v", err)
|
||||
}
|
||||
if doc.Size != int64(len(parts[0])+len(parts[1])+len(parts[2])) {
|
||||
t.Fatalf("doc size = %d", doc.Size)
|
||||
}
|
||||
if blobs.getUploadPartCalls != 0 {
|
||||
t.Fatalf("streaming document path called GetUploadPart %d times", blobs.getUploadPartCalls)
|
||||
}
|
||||
if remaining, err := media.LoadFileParts(ctx, 10, 200); err != nil || len(remaining) != 0 {
|
||||
t.Fatalf("upload parts after success = %+v err=%v", remaining, err)
|
||||
}
|
||||
blob, ok, err := media.GetFileBlob(ctx, "doc:"+strconv.FormatInt(doc.ID, 10))
|
||||
if err != nil || !ok {
|
||||
t.Fatalf("body file blob ok=%v err=%v", ok, err)
|
||||
}
|
||||
body, err := local.Get(ctx, blob.ObjectKey)
|
||||
if err != nil {
|
||||
t.Fatalf("read body blob: %v", err)
|
||||
}
|
||||
if string(body) != strings.Join(parts, "") {
|
||||
t.Fatalf("body blob mismatch")
|
||||
}
|
||||
}
|
||||
|
||||
type countingUploadPartBackend struct {
|
||||
*LocalFS
|
||||
getUploadPartCalls int
|
||||
}
|
||||
|
||||
func (c *countingUploadPartBackend) GetUploadPart(ctx context.Context, objectKey string) ([]byte, error) {
|
||||
c.getUploadPartCalls++
|
||||
return c.LocalFS.GetUploadPart(ctx, objectKey)
|
||||
}
|
||||
|
||||
func newUploadPartTestService(t *testing.T, media *fakeMediaStore, quota domain.UploadPartQuota) (*Service, *LocalFS) {
|
||||
t.Helper()
|
||||
blobs, err := NewLocalFS(t.TempDir())
|
||||
if err != nil {
|
||||
t.Fatalf("NewLocalFS: %v", err)
|
||||
}
|
||||
return NewService(media, blobs, 2,
|
||||
WithVideoThumbnailer(nil),
|
||||
WithUploadPartQuota(quota),
|
||||
), blobs
|
||||
}
|
||||
132
internal/app/files/video_thumbnail.go
Normal file
132
internal/app/files/video_thumbnail.go
Normal file
|
|
@ -0,0 +1,132 @@
|
|||
package files
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"os"
|
||||
"os/exec"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
const (
|
||||
videoThumbnailTimeout = 5 * time.Second
|
||||
videoThumbnailMaxInputBytes = 200 << 20 // 200MB;更大文件本阶段跳过 fallback,避免阻塞发送。
|
||||
videoThumbnailMaxConcurrent = 2
|
||||
)
|
||||
|
||||
// VideoThumbnailer 从视频字节中抽取静态缩略图。实现必须可失败降级,不影响原发送流程。
|
||||
type VideoThumbnailer interface {
|
||||
Extract(ctx context.Context, data []byte, mimeType string) ([]byte, error)
|
||||
}
|
||||
|
||||
// FFmpegVideoThumbnailer 使用本机 ffmpeg 抽取第一帧 JPEG。
|
||||
type FFmpegVideoThumbnailer struct {
|
||||
path string
|
||||
timeout time.Duration
|
||||
slots chan struct{}
|
||||
}
|
||||
|
||||
// NewFFmpegVideoThumbnailer 返回基于 PATH 中 ffmpeg 的抽帧器。
|
||||
func NewFFmpegVideoThumbnailer() (*FFmpegVideoThumbnailer, error) {
|
||||
path, err := exec.LookPath("ffmpeg")
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &FFmpegVideoThumbnailer{
|
||||
path: path,
|
||||
timeout: videoThumbnailTimeout,
|
||||
slots: make(chan struct{}, videoThumbnailMaxConcurrent),
|
||||
}, nil
|
||||
}
|
||||
|
||||
// Extract 抽取第一帧并输出 JPEG bytes。
|
||||
func (t *FFmpegVideoThumbnailer) Extract(ctx context.Context, data []byte, mimeType string) ([]byte, error) {
|
||||
if t == nil || t.path == "" {
|
||||
return nil, fmt.Errorf("ffmpeg thumbnailer unavailable")
|
||||
}
|
||||
if len(data) == 0 {
|
||||
return nil, fmt.Errorf("empty video data")
|
||||
}
|
||||
if len(data) > videoThumbnailMaxInputBytes {
|
||||
return nil, fmt.Errorf("video too large for thumbnail fallback: %d bytes", len(data))
|
||||
}
|
||||
select {
|
||||
case t.slots <- struct{}{}:
|
||||
defer func() { <-t.slots }()
|
||||
case <-ctx.Done():
|
||||
return nil, ctx.Err()
|
||||
}
|
||||
|
||||
runCtx, cancel := context.WithTimeout(ctx, t.timeout)
|
||||
defer cancel()
|
||||
|
||||
input, err := os.CreateTemp("", "telesrv-video-*"+videoTempExt(mimeType))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("create temp video: %w", err)
|
||||
}
|
||||
inputPath := input.Name()
|
||||
defer os.Remove(inputPath)
|
||||
if _, err := input.Write(data); err != nil {
|
||||
input.Close()
|
||||
return nil, fmt.Errorf("write temp video: %w", err)
|
||||
}
|
||||
if err := input.Close(); err != nil {
|
||||
return nil, fmt.Errorf("close temp video: %w", err)
|
||||
}
|
||||
|
||||
output, err := os.CreateTemp("", "telesrv-video-thumb-*.jpg")
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("create temp thumbnail: %w", err)
|
||||
}
|
||||
outputPath := output.Name()
|
||||
output.Close()
|
||||
defer os.Remove(outputPath)
|
||||
|
||||
cmd := exec.CommandContext(
|
||||
runCtx,
|
||||
t.path,
|
||||
"-hide_banner",
|
||||
"-loglevel", "error",
|
||||
"-y",
|
||||
"-i", inputPath,
|
||||
"-map", "0:v:0",
|
||||
"-frames:v", "1",
|
||||
"-an",
|
||||
"-vf", "scale=320:320:force_original_aspect_ratio=decrease",
|
||||
"-q:v", "3",
|
||||
outputPath,
|
||||
)
|
||||
stderr, err := cmd.CombinedOutput()
|
||||
if runCtx.Err() != nil {
|
||||
return nil, runCtx.Err()
|
||||
}
|
||||
if err != nil {
|
||||
msg := strings.TrimSpace(string(stderr))
|
||||
if msg != "" {
|
||||
return nil, fmt.Errorf("ffmpeg extract thumbnail: %w: %s", err, msg)
|
||||
}
|
||||
return nil, fmt.Errorf("ffmpeg extract thumbnail: %w", err)
|
||||
}
|
||||
thumb, err := os.ReadFile(outputPath)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("read thumbnail: %w", err)
|
||||
}
|
||||
if len(thumb) == 0 {
|
||||
return nil, fmt.Errorf("ffmpeg produced empty thumbnail")
|
||||
}
|
||||
return thumb, nil
|
||||
}
|
||||
|
||||
func videoTempExt(mimeType string) string {
|
||||
switch strings.ToLower(strings.TrimSpace(mimeType)) {
|
||||
case "video/mp4":
|
||||
return ".mp4"
|
||||
case "video/quicktime":
|
||||
return ".mov"
|
||||
case "video/webm":
|
||||
return ".webm"
|
||||
default:
|
||||
return ".bin"
|
||||
}
|
||||
}
|
||||
|
|
@ -18,7 +18,19 @@ type WarmStats struct {
|
|||
// SeedMedia 在已有数据时会跳过导入;该方法保证普通 server 重启后历史 sticker 首次渲染也不是冷缓存。
|
||||
func (s *Service) WarmCaches(ctx context.Context) (WarmStats, error) {
|
||||
var stats WarmStats
|
||||
// 第一阶段:收集所有待预热文档(贴纸集 + reaction),按 doc ID 去重。
|
||||
seenDocs := make(map[int64]struct{})
|
||||
docs := make([]domain.Document, 0, 256)
|
||||
collect := func(doc domain.Document) {
|
||||
if doc.ID == 0 {
|
||||
return
|
||||
}
|
||||
if _, ok := seenDocs[doc.ID]; ok {
|
||||
return
|
||||
}
|
||||
seenDocs[doc.ID] = struct{}{}
|
||||
docs = append(docs, doc)
|
||||
}
|
||||
for _, kind := range []domain.StickerSetKind{
|
||||
domain.StickerSetKindStickers,
|
||||
domain.StickerSetKindEmoji,
|
||||
|
|
@ -30,24 +42,15 @@ func (s *Service) WarmCaches(ctx context.Context) (WarmStats, error) {
|
|||
return stats, err
|
||||
}
|
||||
for _, set := range sets {
|
||||
docs, err := s.media.GetDocuments(ctx, set.DocumentIDs)
|
||||
setDocs, err := s.media.GetDocuments(ctx, set.DocumentIDs)
|
||||
if err != nil {
|
||||
return stats, err
|
||||
}
|
||||
ordered := orderDocuments(docs, set.DocumentIDs)
|
||||
ordered := orderDocuments(setDocs, set.DocumentIDs)
|
||||
s.stickerSetCache.put(set, ordered)
|
||||
stats.StickerSets++
|
||||
for _, doc := range ordered {
|
||||
if _, ok := seenDocs[doc.ID]; ok {
|
||||
continue
|
||||
}
|
||||
seenDocs[doc.ID] = struct{}{}
|
||||
stats.Documents++
|
||||
warmed, err := s.prewarmDocumentBlobs(ctx, doc)
|
||||
if err != nil {
|
||||
return stats, err
|
||||
}
|
||||
stats.Blobs += warmed
|
||||
collect(doc)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -59,68 +62,62 @@ func (s *Service) WarmCaches(ctx context.Context) (WarmStats, error) {
|
|||
for _, reaction := range reactions {
|
||||
reactionIDs = append(reactionIDs, reaction.DocumentIDs()...)
|
||||
}
|
||||
docs, err := s.media.GetDocuments(ctx, reactionIDs)
|
||||
reactionDocs, err := s.media.GetDocuments(ctx, reactionIDs)
|
||||
if err != nil {
|
||||
return stats, err
|
||||
}
|
||||
for _, doc := range reactionDocs {
|
||||
collect(doc)
|
||||
}
|
||||
stats.Documents = len(docs)
|
||||
|
||||
// 第二阶段:一发 ANY 查询批量取所有 location key 的 blob 元数据,替代过去逐个
|
||||
// GetFileBlob 的启动期 N+1(~2400 个 blob 各打一次 PG → 一次往返)。
|
||||
keys := make([]string, 0, len(docs)*2)
|
||||
for _, doc := range docs {
|
||||
if _, ok := seenDocs[doc.ID]; ok {
|
||||
keys = append(keys, blobLocationKeys(doc)...)
|
||||
}
|
||||
blobs, err := s.media.GetFileBlobs(ctx, keys)
|
||||
if err != nil {
|
||||
return stats, err
|
||||
}
|
||||
// 第三阶段:填充元数据缓存,并把小 blob 的全量字节读入字节缓存(blob backend 读,非 PG)。
|
||||
for _, key := range keys {
|
||||
blob, ok := blobs[key]
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
seenDocs[doc.ID] = struct{}{}
|
||||
stats.Documents++
|
||||
warmed, err := s.prewarmDocumentBlobs(ctx, doc)
|
||||
s.blobCache.put(key, blob)
|
||||
warmed, err := s.warmBlobBytes(ctx, blob)
|
||||
if err != nil {
|
||||
return stats, err
|
||||
}
|
||||
stats.Blobs += warmed
|
||||
if warmed {
|
||||
stats.Blobs++
|
||||
}
|
||||
}
|
||||
return stats, nil
|
||||
}
|
||||
|
||||
func (s *Service) prewarmDocumentBlobs(ctx context.Context, doc domain.Document) (int, error) {
|
||||
// blobLocationKeys 返回一个文档需预热的全部 location key(主体 + 可下载缩略图)。
|
||||
func blobLocationKeys(doc domain.Document) []string {
|
||||
if doc.ID == 0 {
|
||||
return 0, nil
|
||||
}
|
||||
warmed := 0
|
||||
ok, err := s.prewarmLocationKey(ctx, fmt.Sprintf("doc:%d", doc.ID))
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
if ok {
|
||||
warmed++
|
||||
return nil
|
||||
}
|
||||
keys := make([]string, 0, 1+len(doc.Thumbs))
|
||||
keys = append(keys, fmt.Sprintf("doc:%d", doc.ID))
|
||||
for _, thumb := range doc.Thumbs {
|
||||
if !thumb.Downloadable() {
|
||||
continue
|
||||
}
|
||||
ok, err := s.prewarmLocationKey(ctx, fmt.Sprintf("doc:%d:%s", doc.ID, thumb.Type))
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
if ok {
|
||||
warmed++
|
||||
}
|
||||
keys = append(keys, fmt.Sprintf("doc:%d:%s", doc.ID, thumb.Type))
|
||||
}
|
||||
return warmed, nil
|
||||
return keys
|
||||
}
|
||||
|
||||
func (s *Service) prewarmLocationKey(ctx context.Context, locationKey string) (bool, error) {
|
||||
blob, ok := s.blobCache.get(locationKey)
|
||||
if !ok {
|
||||
var (
|
||||
found bool
|
||||
err error
|
||||
)
|
||||
blob, found, err = s.media.GetFileBlob(ctx, locationKey)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
if !found {
|
||||
return false, nil
|
||||
}
|
||||
s.blobCache.put(locationKey, blob)
|
||||
}
|
||||
// warmBlobBytes 把小 blob 的全量字节读入 byteCache(大 blob 跳过,仍由 GetRange 分段读)。
|
||||
// 返回是否实际写入了字节缓存。
|
||||
func (s *Service) warmBlobBytes(ctx context.Context, blob domain.FileBlob) (bool, error) {
|
||||
if blob.Size <= 0 || blob.Size > blobBytesCacheMaxEntryBytes || s.byteCache.has(blob.ObjectKey) {
|
||||
return false, nil
|
||||
}
|
||||
|
|
|
|||
525
internal/app/files/webpage.go
Normal file
525
internal/app/files/webpage.go
Normal file
|
|
@ -0,0 +1,525 @@
|
|||
package files
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"hash/fnv"
|
||||
"image"
|
||||
"io"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strings"
|
||||
"sync"
|
||||
"syscall"
|
||||
"time"
|
||||
|
||||
"go.uber.org/zap"
|
||||
"golang.org/x/net/html"
|
||||
|
||||
"telesrv/internal/domain"
|
||||
"telesrv/internal/readmodelcache"
|
||||
)
|
||||
|
||||
// 链接预览(webpage preview):抓取消息里的 URL,解析 OpenGraph/Twitter-card/<title>+meta
|
||||
// 元数据,铸造预览卡片(含可选预览图)。安全模型与 external_media 同构(SSRF 拨号期 IP 校验、
|
||||
// 仅 http/https、重定向/大小/超时上限、全局限速),但用独立的限速器与缓存,避免与外链媒体抓取
|
||||
// 争用同一预算。HTML 与预览图共用一个总时长预算(父 ctx deadline)。
|
||||
//
|
||||
// 解析结果(done / empty)经 L1 进程内缓存(singleflight 折叠并发同 URL 抓取)+ L3 web_pages
|
||||
// 表(按规范化 URL 哈希跨实例去重)。瞬时失败(网络/限速/SSRF 拦截)返回 error 不缓存,避免一次
|
||||
// 抖动把热门链接毒成"无预览"。
|
||||
|
||||
var (
|
||||
// ErrWebPagePreviewDisabled 表示未启用链接预览抓取。
|
||||
ErrWebPagePreviewDisabled = errors.New("web page preview disabled")
|
||||
// ErrWebPagePreviewInvalid 表示 URL 不合法/被 SSRF 拦截/上游失败/超限。
|
||||
ErrWebPagePreviewInvalid = errors.New("web page preview invalid")
|
||||
// errWebPageTerminal 标记「确定性、短期不会变」的失败(SSRF 拦截/4xx/非法 URL)。这类
|
||||
// 解析为终态空预览并负缓存,避免每次按键/发送重复打 PG+外网;瞬时失败(5xx/超时/限速/
|
||||
// dial 失败)不带此标记、不缓存、可重试。
|
||||
errWebPageTerminal = errors.New("web page terminal")
|
||||
)
|
||||
|
||||
// terminalFetchErr 构造一个终态失败错误(会被负缓存)。
|
||||
func terminalFetchErr(msg string) error {
|
||||
return fmt.Errorf("%w: %w: %s", ErrWebPagePreviewInvalid, errWebPageTerminal, msg)
|
||||
}
|
||||
|
||||
const (
|
||||
webpageRequestTimeout = 15 * time.Second
|
||||
webpageTotalTimeout = 20 * time.Second
|
||||
webpageMaxRedirects = 5
|
||||
// DefaultWebPagePreviewMaxBytes 覆盖 HTML 抓取与预览图抓取(head 在页首,足够)。
|
||||
DefaultWebPagePreviewMaxBytes = int64(5 << 20)
|
||||
// DefaultWebPagePreviewRatePerMin 是全局每分钟抓取上限;一次解析最多 2 次上游(HTML+图)。
|
||||
// 60 是单用户口径,多用户实例偏低(输入预览与真实发送共用此预算易互相饿死),上调到 300。
|
||||
DefaultWebPagePreviewRatePerMin = 300
|
||||
webpageRateWindow = time.Minute
|
||||
// maxWebpageImagePixels 是预览图解压炸弹上界(解码前按 DecodeConfig 尺寸拦截)。
|
||||
maxWebpageImagePixels = int64(25_000_000)
|
||||
webpageCacheMaxEntries = 4096
|
||||
webpageCacheTTL = 10 * time.Minute
|
||||
// webPageRefreshTTL 是已解析卡片的陈旧阈值:L3 命中且超过此龄时后台 stale-while-revalidate
|
||||
// 刷新(返回的仍是旧卡片,不阻塞)。webPageRefreshConcurrency 限并发刷新 goroutine。
|
||||
webPageRefreshTTL = 24 * time.Hour
|
||||
webPageRefreshConcurrency = 8
|
||||
// webpageUserAgent 用 Telegram 爬虫标识:很多站点只对已知爬虫吐 OG 标签。
|
||||
webpageUserAgent = "TelegramBot (like TwitterBot)"
|
||||
acceptHTML = "text/html,application/xhtml+xml"
|
||||
acceptImage = "image/*"
|
||||
)
|
||||
|
||||
type webpageFetcher struct {
|
||||
client *http.Client
|
||||
maxBytes int64
|
||||
rateLimit int
|
||||
cache *readmodelcache.Cache[int64, domain.MessageWebPage]
|
||||
refreshSem chan struct{}
|
||||
|
||||
mu sync.Mutex
|
||||
fetchTimes []time.Time
|
||||
}
|
||||
|
||||
// WithWebPagePreview 启用链接预览抓取。maxBytes<=0 / ratePerMin<=0 用默认。SSRF 防护恒开。
|
||||
func WithWebPagePreview(maxBytes int64, ratePerMin int) Option {
|
||||
return func(s *Service) {
|
||||
if maxBytes <= 0 {
|
||||
maxBytes = DefaultWebPagePreviewMaxBytes
|
||||
}
|
||||
if ratePerMin <= 0 {
|
||||
ratePerMin = DefaultWebPagePreviewRatePerMin
|
||||
}
|
||||
s.webpage = newWebpageFetcher(maxBytes, ratePerMin, false)
|
||||
}
|
||||
}
|
||||
|
||||
// newWebpageFetcher 构造抓取器。allowPrivate 仅供测试(指向 httptest loopback);生产恒 false。
|
||||
func newWebpageFetcher(maxBytes int64, ratePerMin int, allowPrivate bool) *webpageFetcher {
|
||||
dialer := &net.Dialer{Timeout: webpageRequestTimeout}
|
||||
dialer.Control = func(_, address string, _ syscall.RawConn) error {
|
||||
host, _, err := net.SplitHostPort(address)
|
||||
if err != nil {
|
||||
return ErrWebPagePreviewInvalid
|
||||
}
|
||||
ip := net.ParseIP(host)
|
||||
if ip == nil {
|
||||
return ErrWebPagePreviewInvalid
|
||||
}
|
||||
if !allowPrivate && isBlockedExternalIP(ip) {
|
||||
// SSRF 拦截是确定性失败 → 标记终态供负缓存(否则每次按键重打 PG+重 dial)。
|
||||
return terminalFetchErr("blocked address " + host + " (SSRF guard)")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
client := &http.Client{
|
||||
Timeout: webpageRequestTimeout,
|
||||
Transport: &http.Transport{DialContext: dialer.DialContext, DisableKeepAlives: true, Proxy: nil},
|
||||
CheckRedirect: func(req *http.Request, via []*http.Request) error {
|
||||
if len(via) >= webpageMaxRedirects {
|
||||
return fmt.Errorf("%w: too many redirects", ErrWebPagePreviewInvalid)
|
||||
}
|
||||
if req.URL.Scheme != "http" && req.URL.Scheme != "https" {
|
||||
return fmt.Errorf("%w: blocked redirect scheme %q", ErrWebPagePreviewInvalid, req.URL.Scheme)
|
||||
}
|
||||
return nil
|
||||
},
|
||||
}
|
||||
return &webpageFetcher{
|
||||
client: client,
|
||||
maxBytes: maxBytes,
|
||||
rateLimit: ratePerMin,
|
||||
refreshSem: make(chan struct{}, webPageRefreshConcurrency),
|
||||
cache: readmodelcache.New(readmodelcache.Config[int64, domain.MessageWebPage]{
|
||||
MaxEntries: webpageCacheMaxEntries,
|
||||
TTL: webpageCacheTTL,
|
||||
}),
|
||||
}
|
||||
}
|
||||
|
||||
func (f *webpageFetcher) allowFetch() bool {
|
||||
f.mu.Lock()
|
||||
defer f.mu.Unlock()
|
||||
now := time.Now()
|
||||
kept := f.fetchTimes[:0]
|
||||
for _, at := range f.fetchTimes {
|
||||
if now.Sub(at) <= webpageRateWindow {
|
||||
kept = append(kept, at)
|
||||
}
|
||||
}
|
||||
f.fetchTimes = kept
|
||||
if len(f.fetchTimes) >= f.rateLimit {
|
||||
return false
|
||||
}
|
||||
f.fetchTimes = append(f.fetchTimes, now)
|
||||
return true
|
||||
}
|
||||
|
||||
// fetch 抓取 URL,返回 (字节, content-type)。SSRF 检查在 dial 阶段发生;ctx 承载共享总预算。
|
||||
func (f *webpageFetcher) fetch(ctx context.Context, rawURL, accept string) ([]byte, string, error) {
|
||||
u, err := url.Parse(strings.TrimSpace(rawURL))
|
||||
if err != nil || (u.Scheme != "http" && u.Scheme != "https") || u.Host == "" {
|
||||
return nil, "", terminalFetchErr("bad url") // 非法 URL 是终态
|
||||
}
|
||||
if !f.allowFetch() {
|
||||
return nil, "", fmt.Errorf("%w: rate limited", ErrWebPagePreviewInvalid) // 限速=瞬时,可重试
|
||||
}
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, u.String(), nil)
|
||||
if err != nil {
|
||||
return nil, "", terminalFetchErr("bad request")
|
||||
}
|
||||
req.Header.Set("User-Agent", webpageUserAgent)
|
||||
req.Header.Set("Accept", accept)
|
||||
resp, err := f.client.Do(req)
|
||||
if err != nil {
|
||||
// SSRF 拦截(dial Control 返回的 terminal)经 url.Error 传上来,errors.Is 仍能识别;
|
||||
// 其余 dial/超时错误是瞬时。
|
||||
return nil, "", fmt.Errorf("%w: %v", ErrWebPagePreviewInvalid, err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
// 4xx=确定性(404/403/410…)终态负缓存;5xx=瞬时可重试。
|
||||
if resp.StatusCode >= 400 && resp.StatusCode < 500 {
|
||||
return nil, "", terminalFetchErr(fmt.Sprintf("upstream status %d", resp.StatusCode))
|
||||
}
|
||||
return nil, "", fmt.Errorf("%w: upstream status %d", ErrWebPagePreviewInvalid, resp.StatusCode)
|
||||
}
|
||||
data, err := io.ReadAll(io.LimitReader(resp.Body, f.maxBytes+1))
|
||||
if err != nil {
|
||||
return nil, "", fmt.Errorf("%w: read body: %v", ErrWebPagePreviewInvalid, err)
|
||||
}
|
||||
if len(data) == 0 || int64(len(data)) > f.maxBytes {
|
||||
return nil, "", fmt.Errorf("%w: body size %d", ErrWebPagePreviewInvalid, len(data))
|
||||
}
|
||||
contentType := resp.Header.Get("Content-Type")
|
||||
if i := strings.IndexByte(contentType, ';'); i >= 0 {
|
||||
contentType = contentType[:i]
|
||||
}
|
||||
return data, strings.TrimSpace(strings.ToLower(contentType)), nil
|
||||
}
|
||||
|
||||
// WebPagePreviewEnabled 报告链接预览抓取是否启用。
|
||||
func (s *Service) WebPagePreviewEnabled() bool {
|
||||
return s != nil && s.webpage != nil
|
||||
}
|
||||
|
||||
// LookupWebPage 仅查缓存(L1 进程内 → L3 web_pages 表)返回已解析的链接预览,不抓取。命中
|
||||
// 返回 (page,true);未缓存或未启用返回 false。发送路径用它在 echo 直接带 done 卡片。
|
||||
// 先 Peek L1:客户端输入时 getWebPagePreview 多半已把同一 URL 解析进 L1,发送时即免一次 PG。
|
||||
func (s *Service) LookupWebPage(ctx context.Context, rawURL string) (domain.MessageWebPage, bool) {
|
||||
if s == nil || s.webpage == nil {
|
||||
return domain.MessageWebPage{}, false
|
||||
}
|
||||
normalized, ok := domain.NormalizeWebPageURL(rawURL)
|
||||
if !ok {
|
||||
return domain.MessageWebPage{}, false
|
||||
}
|
||||
urlHash := domain.WebPageURLHash(normalized)
|
||||
if page, ok := s.webpage.cache.Peek(urlHash); ok {
|
||||
return page, true
|
||||
}
|
||||
page, _, found, err := s.media.GetWebPageByURLHash(ctx, urlHash)
|
||||
if err != nil || !found {
|
||||
return domain.MessageWebPage{}, false
|
||||
}
|
||||
s.webpage.cache.Store(urlHash, page) // 回填 L1,后续 Peek 命中。
|
||||
return page, true
|
||||
}
|
||||
|
||||
// ResolveWebPage 解析链接预览,经 L1 缓存(singleflight 去重)+ L3 web_pages 持久去重。
|
||||
// 返回 done / empty 形态的 MessageWebPage;瞬时失败返回 error(调用方降级为空,不报错给用户)。
|
||||
func (s *Service) ResolveWebPage(ctx context.Context, rawURL string) (domain.MessageWebPage, error) {
|
||||
if s == nil || s.webpage == nil {
|
||||
return domain.MessageWebPage{}, ErrWebPagePreviewDisabled
|
||||
}
|
||||
normalized, ok := domain.NormalizeWebPageURL(rawURL)
|
||||
if !ok {
|
||||
return domain.MessageWebPage{}, ErrWebPagePreviewInvalid
|
||||
}
|
||||
urlHash := domain.WebPageURLHash(normalized)
|
||||
return s.webpage.cache.GetOrLoad(ctx, urlHash, func() (domain.MessageWebPage, error) {
|
||||
// L3 durable 命中:直接复用,跨实例去重;超龄则后台刷新(返回旧卡片不阻塞)。
|
||||
if page, refreshedAt, found, err := s.media.GetWebPageByURLHash(ctx, urlHash); err == nil && found {
|
||||
s.webpage.maybeRefresh(s, normalized, urlHash, refreshedAt)
|
||||
return page, nil
|
||||
}
|
||||
// miss:抓取 + 解析(+ 图)。瞬时失败返回 error → GetOrLoad 不缓存(热门链接不被毒化)。
|
||||
page, err := s.webpage.resolve(ctx, s, normalized, urlHash)
|
||||
if err != nil {
|
||||
return domain.MessageWebPage{}, err
|
||||
}
|
||||
// 终态(done / empty):落 L3 供跨实例 + 重启复用。
|
||||
if perr := s.media.PutWebPage(ctx, urlHash, page, int(time.Now().Unix())); perr != nil {
|
||||
s.log.Warn("persist web page preview failed", zap.Int64("url_hash", urlHash), zap.Error(perr))
|
||||
}
|
||||
return page, nil
|
||||
})
|
||||
}
|
||||
|
||||
// maybeRefresh 在卡片超过 webPageRefreshTTL 龄时后台 stale-while-revalidate 刷新一次。
|
||||
// 并发受 refreshSem 限;满则跳过(下次再刷)。瞬时失败保留旧卡片。
|
||||
func (f *webpageFetcher) maybeRefresh(s *Service, normalizedURL string, urlHash int64, refreshedAt int) {
|
||||
if refreshedAt == 0 || time.Now().Unix()-int64(refreshedAt) < int64(webPageRefreshTTL/time.Second) {
|
||||
return
|
||||
}
|
||||
select {
|
||||
case f.refreshSem <- struct{}{}:
|
||||
default:
|
||||
return
|
||||
}
|
||||
go func() {
|
||||
defer func() { <-f.refreshSem }()
|
||||
ctx, cancel := context.WithTimeout(context.Background(), webpageTotalTimeout)
|
||||
defer cancel()
|
||||
page, err := f.resolve(ctx, s, normalizedURL, urlHash)
|
||||
if err != nil {
|
||||
return // 瞬时失败:保留旧卡片。
|
||||
}
|
||||
if perr := s.media.PutWebPage(ctx, urlHash, page, int(time.Now().Unix())); perr != nil {
|
||||
return
|
||||
}
|
||||
f.cache.Store(urlHash, page) // 刷新 L1,使后续读到新卡片。
|
||||
}()
|
||||
}
|
||||
|
||||
// resolve 实际抓取并构造卡片。HTML 与预览图共享 ctx 总时长预算。
|
||||
func (f *webpageFetcher) resolve(ctx context.Context, s *Service, normalizedURL string, urlHash int64) (domain.MessageWebPage, error) {
|
||||
ctx, cancel := context.WithTimeout(ctx, webpageTotalTimeout)
|
||||
defer cancel()
|
||||
|
||||
data, contentType, err := f.fetch(ctx, normalizedURL, acceptHTML)
|
||||
if err != nil {
|
||||
// 终态失败(SSRF/4xx/非法 URL)→ 负缓存为空预览,避免重复按键/发送重打 PG+外网。
|
||||
// 瞬时失败(5xx/超时/dial/限速)→ 上抛 error,GetOrLoad 不缓存、可重试。
|
||||
if errors.Is(err, errWebPageTerminal) {
|
||||
return emptyWebPage(normalizedURL, urlHash), nil
|
||||
}
|
||||
return domain.MessageWebPage{}, err
|
||||
}
|
||||
if !isHTMLContentType(contentType) {
|
||||
// 非 HTML(如直接指向图片/二进制):终态空预览。
|
||||
return emptyWebPage(normalizedURL, urlHash), nil
|
||||
}
|
||||
meta := parseWebPageMeta(data, normalizedURL)
|
||||
if meta.empty() {
|
||||
return emptyWebPage(normalizedURL, urlHash), nil
|
||||
}
|
||||
page := doneWebPage(meta, normalizedURL, urlHash)
|
||||
if meta.image != "" {
|
||||
if photo, ok := f.fetchImage(ctx, s, meta.image); ok {
|
||||
page.Photo = &photo
|
||||
page.HasLargeMedia = true
|
||||
}
|
||||
}
|
||||
return page, nil
|
||||
}
|
||||
|
||||
// fetchImage 抓取并铸造预览图(best-effort)。解码前按尺寸拦截解压炸弹;非图片/失败丢弃。
|
||||
func (f *webpageFetcher) fetchImage(ctx context.Context, s *Service, imageURL string) (domain.Photo, bool) {
|
||||
data, _, err := f.fetch(ctx, imageURL, acceptImage)
|
||||
if err != nil {
|
||||
return domain.Photo{}, false
|
||||
}
|
||||
cfg, _, derr := image.DecodeConfig(bytes.NewReader(data))
|
||||
if derr != nil || cfg.Width <= 0 || cfg.Height <= 0 || int64(cfg.Width)*int64(cfg.Height) > maxWebpageImagePixels {
|
||||
return domain.Photo{}, false
|
||||
}
|
||||
photo, err := s.CreatePhotoFromBytes(ctx, data)
|
||||
if err != nil {
|
||||
return domain.Photo{}, false
|
||||
}
|
||||
return photo, true
|
||||
}
|
||||
|
||||
func emptyWebPage(rawURL string, urlHash int64) domain.MessageWebPage {
|
||||
return domain.MessageWebPage{State: domain.MessageWebPageStateEmpty, ID: urlHash, URL: rawURL}
|
||||
}
|
||||
|
||||
func doneWebPage(meta webpageMeta, rawURL string, urlHash int64) domain.MessageWebPage {
|
||||
page := domain.MessageWebPage{
|
||||
State: domain.MessageWebPageStateDone,
|
||||
ID: urlHash,
|
||||
URL: rawURL,
|
||||
DisplayURL: webpageDisplayURL(rawURL),
|
||||
Type: meta.pageType(),
|
||||
SiteName: meta.siteName,
|
||||
Title: meta.title,
|
||||
Description: meta.description,
|
||||
Author: meta.author,
|
||||
}
|
||||
page.Hash = webpageContentHash(page)
|
||||
return page
|
||||
}
|
||||
|
||||
// webpageMeta 是从 HTML head 提取的预览元数据(已按 og>twitter>title/meta 优先级归并)。
|
||||
type webpageMeta struct {
|
||||
title string
|
||||
description string
|
||||
siteName string
|
||||
image string
|
||||
author string
|
||||
ogType string
|
||||
}
|
||||
|
||||
func (m webpageMeta) empty() bool {
|
||||
return m.title == "" && m.description == "" && m.siteName == "" && m.image == "" && m.author == ""
|
||||
}
|
||||
|
||||
func (m webpageMeta) pageType() string {
|
||||
if m.ogType != "" {
|
||||
return m.ogType
|
||||
}
|
||||
if m.image != "" && m.title == "" && m.description == "" {
|
||||
return "photo"
|
||||
}
|
||||
return "article"
|
||||
}
|
||||
|
||||
// parseWebPageMeta 扫描 HTML head 的 <meta>/<title>,提取 OpenGraph/Twitter-card/标准元数据。
|
||||
// 遇到 <body> 或 </head> 即停止(元数据都在 head)。og:image 相对 URL 按 baseURL 解析为绝对。
|
||||
func parseWebPageMeta(htmlBytes []byte, baseURL string) webpageMeta {
|
||||
var (
|
||||
m webpageMeta
|
||||
ogTitle, twTitle, docTitle string
|
||||
ogDesc, twDesc, metaDesc string
|
||||
ogImage, twImage string
|
||||
inTitle bool
|
||||
)
|
||||
z := html.NewTokenizer(bytes.NewReader(htmlBytes))
|
||||
scan:
|
||||
for {
|
||||
switch z.Next() {
|
||||
case html.ErrorToken:
|
||||
break scan
|
||||
case html.StartTagToken, html.SelfClosingTagToken:
|
||||
name, hasAttr := z.TagName()
|
||||
switch string(name) {
|
||||
case "meta":
|
||||
key, content := metaKeyContent(z, hasAttr)
|
||||
if content == "" {
|
||||
continue
|
||||
}
|
||||
switch key {
|
||||
case "og:title":
|
||||
ogTitle = content
|
||||
case "og:description":
|
||||
ogDesc = content
|
||||
case "og:site_name":
|
||||
m.siteName = content
|
||||
case "og:type":
|
||||
m.ogType = content
|
||||
case "og:image", "og:image:url", "og:image:secure_url":
|
||||
if ogImage == "" {
|
||||
ogImage = content
|
||||
}
|
||||
case "twitter:title":
|
||||
twTitle = content
|
||||
case "twitter:description":
|
||||
twDesc = content
|
||||
case "twitter:image", "twitter:image:src":
|
||||
if twImage == "" {
|
||||
twImage = content
|
||||
}
|
||||
case "description":
|
||||
metaDesc = content
|
||||
case "author", "article:author":
|
||||
if m.author == "" {
|
||||
m.author = content
|
||||
}
|
||||
}
|
||||
case "title":
|
||||
inTitle = true
|
||||
case "body":
|
||||
break scan
|
||||
}
|
||||
case html.TextToken:
|
||||
if inTitle && docTitle == "" {
|
||||
docTitle = strings.TrimSpace(string(z.Text()))
|
||||
}
|
||||
case html.EndTagToken:
|
||||
name, _ := z.TagName()
|
||||
switch string(name) {
|
||||
case "title":
|
||||
inTitle = false
|
||||
case "head":
|
||||
break scan
|
||||
}
|
||||
}
|
||||
}
|
||||
m.title = firstNonEmpty(ogTitle, twTitle, docTitle)
|
||||
m.description = firstNonEmpty(ogDesc, twDesc, metaDesc)
|
||||
if img := firstNonEmpty(ogImage, twImage); img != "" {
|
||||
if abs, ok := resolveAbsoluteURL(baseURL, img); ok {
|
||||
m.image = abs
|
||||
}
|
||||
}
|
||||
return m
|
||||
}
|
||||
|
||||
// metaKeyContent 从一个 <meta> 标签收集 (property|name) 与 content。
|
||||
func metaKeyContent(z *html.Tokenizer, hasAttr bool) (string, string) {
|
||||
var key, content string
|
||||
for hasAttr {
|
||||
var k, v []byte
|
||||
k, v, hasAttr = z.TagAttr()
|
||||
switch strings.ToLower(string(k)) {
|
||||
case "property", "name":
|
||||
if key == "" {
|
||||
key = strings.ToLower(strings.TrimSpace(string(v)))
|
||||
}
|
||||
case "content":
|
||||
content = strings.TrimSpace(string(v))
|
||||
}
|
||||
}
|
||||
return key, content
|
||||
}
|
||||
|
||||
func resolveAbsoluteURL(base, ref string) (string, bool) {
|
||||
b, err := url.Parse(base)
|
||||
if err != nil {
|
||||
return "", false
|
||||
}
|
||||
r, err := url.Parse(strings.TrimSpace(ref))
|
||||
if err != nil {
|
||||
return "", false
|
||||
}
|
||||
abs := b.ResolveReference(r)
|
||||
if abs.Scheme != "http" && abs.Scheme != "https" {
|
||||
return "", false
|
||||
}
|
||||
return abs.String(), true
|
||||
}
|
||||
|
||||
func webpageDisplayURL(raw string) string {
|
||||
u, err := url.Parse(raw)
|
||||
if err != nil || u.Host == "" {
|
||||
return raw
|
||||
}
|
||||
return strings.TrimPrefix(u.Host, "www.")
|
||||
}
|
||||
|
||||
// webpageContentHash 对卡片内容算稳定 31-bit 哈希(webPage.hash 是 TL int,用于 getWebPage
|
||||
// NotModified 短路)。仅覆盖文本字段,预览图变化不计入(同 URL 预览图按内容寻址已去重)。
|
||||
func webpageContentHash(p domain.MessageWebPage) int {
|
||||
h := fnv.New32a()
|
||||
for _, s := range []string{p.URL, p.Title, p.Description, p.SiteName, p.Author, p.Type} {
|
||||
_, _ = h.Write([]byte(s))
|
||||
_, _ = h.Write([]byte{0})
|
||||
}
|
||||
return int(h.Sum32() & 0x7fffffff)
|
||||
}
|
||||
|
||||
func isHTMLContentType(ct string) bool {
|
||||
return ct == "" || ct == "text/html" || ct == "application/xhtml+xml"
|
||||
}
|
||||
|
||||
func firstNonEmpty(vals ...string) string {
|
||||
for _, v := range vals {
|
||||
if v != "" {
|
||||
return v
|
||||
}
|
||||
}
|
||||
return ""
|
||||
}
|
||||
228
internal/app/files/webpage_test.go
Normal file
228
internal/app/files/webpage_test.go
Normal file
|
|
@ -0,0 +1,228 @@
|
|||
package files
|
||||
|
||||
import (
|
||||
"context"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"telesrv/internal/domain"
|
||||
)
|
||||
|
||||
// newWebpageTestService 构造一个带 loopback-allowed 抓取器的 Service(生产恒禁 loopback)。
|
||||
func newWebpageTestService(t *testing.T, allowPrivate bool) *Service {
|
||||
t.Helper()
|
||||
media := newFakeMediaStore()
|
||||
blobs, err := NewLocalFS(t.TempDir())
|
||||
if err != nil {
|
||||
t.Fatalf("NewLocalFS: %v", err)
|
||||
}
|
||||
svc := NewService(media, blobs, 2)
|
||||
svc.webpage = newWebpageFetcher(DefaultWebPagePreviewMaxBytes, 600, allowPrivate)
|
||||
return svc
|
||||
}
|
||||
|
||||
func TestResolveWebPageDoneCardWithImage(t *testing.T) {
|
||||
imgBytes := testJPEG(t, 8, 6)
|
||||
var pageHits, imgHits int32
|
||||
mux := http.NewServeMux()
|
||||
mux.HandleFunc("/img.jpg", func(w http.ResponseWriter, _ *http.Request) {
|
||||
atomic.AddInt32(&imgHits, 1)
|
||||
w.Header().Set("Content-Type", "image/jpeg")
|
||||
_, _ = w.Write(imgBytes)
|
||||
})
|
||||
mux.HandleFunc("/article", func(w http.ResponseWriter, _ *http.Request) {
|
||||
atomic.AddInt32(&pageHits, 1)
|
||||
w.Header().Set("Content-Type", "text/html; charset=utf-8")
|
||||
_, _ = io.WriteString(w, `<html><head>
|
||||
<title>Fallback Title</title>
|
||||
<meta property="og:title" content="OG Title">
|
||||
<meta property="og:description" content="OG Description">
|
||||
<meta property="og:site_name" content="Example Site">
|
||||
<meta property="og:type" content="article">
|
||||
<meta property="og:image" content="/img.jpg">
|
||||
</head><body>ignored body</body></html>`)
|
||||
})
|
||||
srv := httptest.NewServer(mux)
|
||||
defer srv.Close()
|
||||
|
||||
svc := newWebpageTestService(t, true)
|
||||
ctx := context.Background()
|
||||
pageURL := srv.URL + "/article"
|
||||
|
||||
page, err := svc.ResolveWebPage(ctx, pageURL)
|
||||
if err != nil {
|
||||
t.Fatalf("ResolveWebPage: %v", err)
|
||||
}
|
||||
if page.State != domain.MessageWebPageStateDone {
|
||||
t.Fatalf("state = %q, want done", page.State)
|
||||
}
|
||||
if page.Title != "OG Title" || page.Description != "OG Description" || page.SiteName != "Example Site" || page.Type != "article" {
|
||||
t.Fatalf("card fields = %+v", page)
|
||||
}
|
||||
if page.Photo == nil || page.Photo.ID == 0 || len(page.Photo.Sizes) == 0 {
|
||||
t.Fatalf("expected minted preview photo, got %+v", page.Photo)
|
||||
}
|
||||
// id == url_hash(保证 pending↔done 关联)。
|
||||
normalized, _ := domain.NormalizeWebPageURL(pageURL)
|
||||
if page.ID != domain.WebPageURLHash(normalized) {
|
||||
t.Fatalf("webPage id %d != url_hash %d", page.ID, domain.WebPageURLHash(normalized))
|
||||
}
|
||||
|
||||
// 第二次解析命中 L1 缓存(singleflight/LRU),不再打上游。
|
||||
if _, err := svc.ResolveWebPage(ctx, pageURL); err != nil {
|
||||
t.Fatalf("second ResolveWebPage: %v", err)
|
||||
}
|
||||
if h := atomic.LoadInt32(&pageHits); h != 1 {
|
||||
t.Fatalf("page fetched %d times, want 1 (cache dedup)", h)
|
||||
}
|
||||
if h := atomic.LoadInt32(&imgHits); h != 1 {
|
||||
t.Fatalf("image fetched %d times, want 1", h)
|
||||
}
|
||||
}
|
||||
|
||||
// TestWebPageMaybeRefreshStale 验证超龄卡片触发后台 stale-while-revalidate 刷新(写回 done)。
|
||||
func TestWebPageMaybeRefreshStale(t *testing.T) {
|
||||
mux := http.NewServeMux()
|
||||
mux.HandleFunc("/article", func(w http.ResponseWriter, _ *http.Request) {
|
||||
w.Header().Set("Content-Type", "text/html")
|
||||
_, _ = io.WriteString(w, `<html><head><meta property="og:title" content="Fresh"></head></html>`)
|
||||
})
|
||||
srv := httptest.NewServer(mux)
|
||||
defer srv.Close()
|
||||
|
||||
ctx := context.Background()
|
||||
svc := newWebpageTestService(t, true)
|
||||
normalized, _ := domain.NormalizeWebPageURL(srv.URL + "/article")
|
||||
urlHash := domain.WebPageURLHash(normalized)
|
||||
|
||||
// 触发刷新(refreshedAt=1 → 远超 TTL)。
|
||||
svc.webpage.maybeRefresh(svc, normalized, urlHash, 1)
|
||||
|
||||
// 轮询直到刷新写回 done 卡片。
|
||||
var ok bool
|
||||
for i := 0; i < 100; i++ {
|
||||
if page, _, found, err := svc.media.GetWebPageByURLHash(ctx, urlHash); err == nil && found && page.State == domain.MessageWebPageStateDone && page.Title == "Fresh" {
|
||||
ok = true
|
||||
break
|
||||
}
|
||||
time.Sleep(20 * time.Millisecond)
|
||||
}
|
||||
if !ok {
|
||||
t.Fatalf("stale refresh did not write fresh done card")
|
||||
}
|
||||
|
||||
// refreshedAt=0 或新鲜 → 不刷新(不 panic、立即返回)。
|
||||
svc.webpage.maybeRefresh(svc, normalized, urlHash, 0)
|
||||
svc.webpage.maybeRefresh(svc, normalized, urlHash, int(time.Now().Unix()))
|
||||
}
|
||||
|
||||
// TestResolveWebPageTerminalFailureNegativeCached 验证 4xx(终态)解析为空预览并负缓存——
|
||||
// 第二次不再打上游(否则每次按键/发送重复抓取坏 URL)。
|
||||
func TestResolveWebPageTerminalFailureNegativeCached(t *testing.T) {
|
||||
var hits int32
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
atomic.AddInt32(&hits, 1)
|
||||
w.WriteHeader(http.StatusNotFound)
|
||||
}))
|
||||
defer srv.Close()
|
||||
svc := newWebpageTestService(t, true)
|
||||
ctx := context.Background()
|
||||
|
||||
page, err := svc.ResolveWebPage(ctx, srv.URL+"/x")
|
||||
if err != nil {
|
||||
t.Fatalf("404 should be terminal-empty (not error): %v", err)
|
||||
}
|
||||
if page.State != domain.MessageWebPageStateEmpty {
|
||||
t.Fatalf("state = %q, want empty", page.State)
|
||||
}
|
||||
if _, err := svc.ResolveWebPage(ctx, srv.URL+"/x"); err != nil {
|
||||
t.Fatalf("second resolve: %v", err)
|
||||
}
|
||||
if h := atomic.LoadInt32(&hits); h != 1 {
|
||||
t.Fatalf("terminal URL fetched %d times, want 1 (negative cached)", h)
|
||||
}
|
||||
}
|
||||
|
||||
// TestResolveWebPageTransientFailureNotCached 验证 5xx(瞬时)返回 error 且不缓存——可重试。
|
||||
func TestResolveWebPageTransientFailureNotCached(t *testing.T) {
|
||||
var hits int32
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
atomic.AddInt32(&hits, 1)
|
||||
w.WriteHeader(http.StatusInternalServerError)
|
||||
}))
|
||||
defer srv.Close()
|
||||
svc := newWebpageTestService(t, true)
|
||||
ctx := context.Background()
|
||||
|
||||
if _, err := svc.ResolveWebPage(ctx, srv.URL+"/x"); err == nil {
|
||||
t.Fatalf("500 should return transient error")
|
||||
}
|
||||
_, _ = svc.ResolveWebPage(ctx, srv.URL+"/x")
|
||||
if h := atomic.LoadInt32(&hits); h != 2 {
|
||||
t.Fatalf("transient URL fetched %d times, want 2 (not cached, retryable)", h)
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveWebPageNoMetadataIsEmpty(t *testing.T) {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
w.Header().Set("Content-Type", "text/html")
|
||||
_, _ = io.WriteString(w, `<html><head></head><body>no meta here</body></html>`)
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
svc := newWebpageTestService(t, true)
|
||||
page, err := svc.ResolveWebPage(context.Background(), srv.URL+"/x")
|
||||
if err != nil {
|
||||
t.Fatalf("ResolveWebPage: %v", err)
|
||||
}
|
||||
if page.State != domain.MessageWebPageStateEmpty {
|
||||
t.Fatalf("state = %q, want empty", page.State)
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveWebPageNonHTMLIsEmpty(t *testing.T) {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/pdf")
|
||||
_, _ = w.Write([]byte("%PDF-1.4 binary"))
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
svc := newWebpageTestService(t, true)
|
||||
page, err := svc.ResolveWebPage(context.Background(), srv.URL+"/doc")
|
||||
if err != nil {
|
||||
t.Fatalf("ResolveWebPage: %v", err)
|
||||
}
|
||||
if page.State != domain.MessageWebPageStateEmpty {
|
||||
t.Fatalf("state = %q, want empty for non-HTML", page.State)
|
||||
}
|
||||
}
|
||||
|
||||
// TestResolveWebPageSSRFBlocksLoopback 验证生产配置(allowPrivate=false)拦截指向 loopback 的 URL。
|
||||
func TestResolveWebPageSSRFBlocksLoopback(t *testing.T) {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
w.Header().Set("Content-Type", "text/html")
|
||||
_, _ = io.WriteString(w, `<html><head><meta property="og:title" content="secret"></head></html>`)
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
svc := newWebpageTestService(t, false) // 生产口径:禁 loopback
|
||||
if _, err := svc.ResolveWebPage(context.Background(), srv.URL+"/x"); err == nil {
|
||||
t.Fatalf("expected SSRF guard to block loopback fetch")
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveWebPageDisabled(t *testing.T) {
|
||||
media := newFakeMediaStore()
|
||||
blobs, err := NewLocalFS(t.TempDir())
|
||||
if err != nil {
|
||||
t.Fatalf("NewLocalFS: %v", err)
|
||||
}
|
||||
svc := NewService(media, blobs, 2) // 未启用 webpage 抓取
|
||||
if _, err := svc.ResolveWebPage(context.Background(), "https://example.com"); err == nil {
|
||||
t.Fatalf("expected ErrWebPagePreviewDisabled")
|
||||
}
|
||||
}
|
||||
Loading…
Add table
Add a link
Reference in a new issue