chore: refresh gramsrv public release

This commit is contained in:
A 2026-06-30 14:37:43 +08:00
parent 75cebe8dbf
commit 70b6820474
1274 changed files with 378751 additions and 59919 deletions

View 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
}

View 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)
}
}

View file

@ -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/头像更被大量用户重复拉)。

View file

@ -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
}

View file

@ -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)
}
}

View 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 hashmessages.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)
}

View 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)
}
// ResolveStickerSetinputStickerSetEmojiDefaultStatuses 的服务路径)能解析。
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")
}
}

View 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 字段算稳定正整数 hashFNV-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)
}

View 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
}

View 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 并铸造 PhotoCreatePhotoFromBytes 会解码校验是否为图片)。
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 并铸造 Documentmime 取 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"
}

View 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=falseloopback 目标被 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)
}
}

View 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。
// 入参越界时按协议约束 clampw/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
}

View 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 配额。
//
// 地图不内嵌定位针TDesktophistoryMapPoint 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 静态图只支持 @2xscale 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 ctxsingleflight 结果被并发分片共享,单个调用方取消不应拖垮整次抓取。
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 ""
}
}

View 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")
}
}

View 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")
}
}

View 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、mdatmoov 已在 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
}
// patchStcostco = [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
}
// patchCo64co64 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
}

View 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前面填充 + markerstco 指向 marker 的绝对偏移。
pad := bytes.Repeat([]byte{0xAB}, 40)
mdatPayload := append(append([]byte(nil), pad...), marker...)
_, mdatAbs := buildMoovEndMP4(mdatPayload)
markerAbs := mdatAbs + uint32(len(pad)) // marker 在原文件里的绝对偏移
// 原 stco 指向 mdatAbsmdat 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 truemoov 在末尾应被搬动)")
}
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) 已在 idxpayload 从 idx+4 起 = ver+flags(4)+count(4)+offset(4)
off := idx + 4 + 4 + 4
return binary.BigEndian.Uint32(moov[off : off+4])
}

View file

@ -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
}
// faststartMP4 视频若 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[:])

View 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()
}

View file

@ -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 setsdefault 系统集 + 常规集)
phaseStarted = time.Now()
phaseBefore = stats
// sticker setsdefault 系统集 + 常规集 + 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

View 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
}

View file

@ -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

View file

@ -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

View 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 = starsv1 全额转换,视作用新购 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
}

View 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))
}
}

View 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
}

View 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"
}
}

View file

@ -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
}

View 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/限速)→ 上抛 errorGetOrLoad 不缓存、可重试。
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 ""
}

View 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")
}
}