diff --git a/internal/httpgetter/html_meta.go b/internal/httpgetter/html_meta.go index 8a5f3149..cc179ef0 100644 --- a/internal/httpgetter/html_meta.go +++ b/internal/httpgetter/html_meta.go @@ -8,16 +8,29 @@ import ( "net/http" "net/url" "strings" + "sync" "time" "github.com/pkg/errors" "golang.org/x/net/html" - "golang.org/x/net/html/atom" + "golang.org/x/sync/semaphore" + "golang.org/x/sync/singleflight" ) +// ErrInternalIP indicates that a link preview target resolves to a disallowed internal address. var ErrInternalIP = errors.New("internal IP addresses are not allowed") -const maxHTMLMetaBytes = 512 * 1024 +const ( + defaultLinkPreviewUserAgent = "MemosBot/1.0 (+https://usememos.com)" + + maxHTMLMetaBytes = 512 * 1024 + maxOEmbedBytes = 128 * 1024 + maxCacheEntries = 1000 + maxConcurrentFetches = 8 + fetchTimeout = 5 * time.Second + successCacheTTL = 24 * time.Hour + failureCacheTTL = time.Minute +) var ( lookupIPAddr = net.DefaultResolver.LookupIPAddr @@ -25,7 +38,6 @@ var ( Timeout: 30 * time.Second, KeepAlive: 30 * time.Second, }).DialContext - httpClient = newHTTPClient() ) func newHTTPClient() *http.Client { @@ -35,7 +47,7 @@ func newHTTPClient() *http.Client { return &http.Client{ Transport: transport, - Timeout: 5 * time.Second, + Timeout: fetchTimeout, CheckRedirect: func(req *http.Request, via []*http.Request) error { if err := validateURL(req.URL.String()); err != nil { return errors.Wrap(err, "redirect to internal IP") @@ -127,101 +139,272 @@ func validateURL(urlStr string) error { return nil } +// HTMLMeta contains metadata used to build a link preview. type HTMLMeta struct { Title string `json:"title"` Description string `json:"description"` Image string `json:"image"` } -func GetHTMLMeta(urlStr string) (*HTMLMeta, error) { - if err := validateURL(urlStr); err != nil { - return nil, err - } +type cacheEntry struct { + meta *HTMLMeta + err error + expiresAt time.Time + storedAt time.Time +} - response, err := httpClient.Get(urlStr) +// HTMLMetaFetcher fetches and caches standards-based link preview metadata. +type HTMLMetaFetcher struct { + client *http.Client + semaphore *semaphore.Weighted + group singleflight.Group + now func() time.Time + + cacheMu sync.Mutex + cache map[string]cacheEntry +} + +// NewHTMLMetaFetcher creates a link preview fetcher with an SSRF-safe HTTP client. +func NewHTMLMetaFetcher() *HTMLMetaFetcher { + return &HTMLMetaFetcher{ + client: newHTTPClient(), + semaphore: semaphore.NewWeighted(maxConcurrentFetches), + now: time.Now, + cache: make(map[string]cacheEntry), + } +} + +// Get fetches metadata for a URL. +func (f *HTMLMetaFetcher) Get(ctx context.Context, urlStr string) (*HTMLMeta, error) { + key, err := normalizeURL(urlStr) if err != nil { return nil, err } + if meta, ok, err := f.getCached(key); ok { + return meta, err + } + + resultChannel := f.group.DoChan(key, func() (any, error) { + if meta, ok, err := f.getCached(key); ok { + return meta, err + } + + // The flight is shared by every coalesced waiter, so one caller's + // cancellation must not abort it for the remaining waiters. The single + // deadline covers both queueing and outbound requests. + flightContext, cancel := context.WithTimeout(context.WithoutCancel(ctx), fetchTimeout) + defer cancel() + if err := f.semaphore.Acquire(flightContext, 1); err != nil { + return nil, err + } + defer f.semaphore.Release(1) + + meta, err := f.fetch(flightContext, key) + if err == nil { + f.setCached(key, meta, nil, successCacheTTL) + } else { + f.setCached(key, nil, err, failureCacheTTL) + } + return meta, err + }) + + select { + case <-ctx.Done(): + return nil, ctx.Err() + case result := <-resultChannel: + if result.Err != nil { + return nil, result.Err + } + meta, ok := result.Val.(*HTMLMeta) + if !ok { + return nil, errors.New("invalid link metadata result") + } + return cloneHTMLMeta(meta), nil + } +} + +func (f *HTMLMetaFetcher) fetch(ctx context.Context, urlStr string) (*HTMLMeta, error) { + request, err := http.NewRequestWithContext(ctx, http.MethodGet, urlStr, nil) + if err != nil { + return nil, errors.Wrap(err, "failed to create link preview request") + } + setRequestHeaders(request, "text/html, application/xhtml+xml;q=0.9, */*;q=0.1") + + response, err := f.client.Do(request) + if err != nil { + return nil, errors.Wrap(err, "failed to fetch link preview") + } defer response.Body.Close() + if response.StatusCode < http.StatusOK || response.StatusCode >= http.StatusMultipleChoices { + return nil, errors.Errorf("unexpected HTTP status: %s", response.Status) + } - mediatype, err := getMediatype(response) + mediaType, err := getMediatype(response) if err != nil { return nil, err } - if mediatype != "text/html" { + if mediaType != "text/html" && mediaType != "application/xhtml+xml" { return nil, errors.New("not a HTML page") } - htmlMeta := extractHTMLMeta(io.LimitReader(response.Body, maxHTMLMetaBytes)) - enrichSiteMeta(response.Request.URL, htmlMeta) - return htmlMeta, nil -} - -func extractHTMLMeta(resp io.Reader) *HTMLMeta { - tokenizer := html.NewTokenizer(resp) - htmlMeta := new(HTMLMeta) - - for { - tokenType := tokenizer.Next() - if tokenType == html.ErrorToken { - break - } else if tokenType == html.StartTagToken || tokenType == html.SelfClosingTagToken { - token := tokenizer.Token() - if token.DataAtom == atom.Body { - break - } - - if token.DataAtom == atom.Title { - tokenizer.Next() - token := tokenizer.Token() - htmlMeta.Title = token.Data - } else if token.DataAtom == atom.Meta { - ogTitle, ok := extractMetaProperty(token, "og:title") - if ok { - htmlMeta.Title = ogTitle - } - - ogDescription, ok := extractMetaProperty(token, "og:description") - if ok { - htmlMeta.Description = ogDescription - } - - ogImage, ok := extractMetaProperty(token, "og:image") - if ok { - htmlMeta.Image = ogImage - } - - description, ok := extractMetaProperty(token, "description") - if ok && htmlMeta.Description == "" { - htmlMeta.Description = description - } - } - } + document, err := html.Parse(io.LimitReader(response.Body, maxHTMLMetaBytes)) + if err != nil { + return nil, errors.Wrap(err, "failed to parse HTML") + } + pageURL, err := finalResponseURL(response, urlStr) + if err != nil { + return nil, err } - return htmlMeta -} - -func extractMetaProperty(token html.Token, prop string) (content string, ok bool) { - content, ok = "", false - for _, attr := range token.Attr { - if (attr.Key == "property" || attr.Key == "name") && strings.EqualFold(attr.Val, prop) { - ok = true - } - if attr.Key == "content" { - content = attr.Val - } + sources, oEmbedEndpoint := extractDocumentMetadata(document) + baseURL := resolveDocumentBase(pageURL, sources.baseHref) + var oEmbed metadataSource + if endpoint := resolveHTTPURL(baseURL, oEmbedEndpoint); endpoint != "" { + // oEmbed is optional enrichment, so its failure must not discard the + // metadata already extracted from the HTML document. + oEmbed, _ = f.fetchOEmbed(ctx, endpoint) } - return content, ok + + meta := mergeMetadata(baseURL, oEmbed, sources.openGraph, sources.twitter, sources.jsonLD, sources.standard, siteImageSource(pageURL), sources.semantic) + return meta, nil } -func enrichSiteMeta(url *url.URL, meta *HTMLMeta) { - if url.Hostname() == "www.youtube.com" { - if url.Path == "/watch" { - vid := url.Query().Get("v") - if vid != "" { - meta.Image = fmt.Sprintf("https://img.youtube.com/vi/%s/mqdefault.jpg", vid) +func (f *HTMLMetaFetcher) fetchOEmbed(ctx context.Context, endpoint string) (metadataSource, error) { + if err := validateURL(endpoint); err != nil { + return metadataSource{}, err + } + request, err := http.NewRequestWithContext(ctx, http.MethodGet, endpoint, nil) + if err != nil { + return metadataSource{}, errors.Wrap(err, "failed to create oEmbed request") + } + setRequestHeaders(request, "application/json+oembed, application/json;q=0.9") + + response, err := f.client.Do(request) + if err != nil { + return metadataSource{}, errors.Wrap(err, "failed to fetch oEmbed endpoint") + } + defer response.Body.Close() + if response.StatusCode < http.StatusOK || response.StatusCode >= http.StatusMultipleChoices { + return metadataSource{}, errors.Errorf("unexpected oEmbed HTTP status: %s", response.Status) + } + + mediaType, err := getMediatype(response) + if err != nil { + return metadataSource{}, err + } + if mediaType != "application/json" && mediaType != "application/json+oembed" && !strings.HasSuffix(mediaType, "+json") { + return metadataSource{}, errors.New("oEmbed response is not JSON") + } + + data, err := io.ReadAll(io.LimitReader(response.Body, maxOEmbedBytes+1)) + if err != nil { + return metadataSource{}, errors.Wrap(err, "failed to read oEmbed response") + } + if len(data) > maxOEmbedBytes { + return metadataSource{}, errors.New("oEmbed response exceeds size limit") + } + + baseURL, err := finalResponseURL(response, endpoint) + if err != nil { + return metadataSource{}, err + } + return extractOEmbedMetadata(data, baseURL) +} + +func setRequestHeaders(request *http.Request, accept string) { + request.Header.Set("User-Agent", defaultLinkPreviewUserAgent) + request.Header.Set("Accept", accept) +} + +func (f *HTMLMetaFetcher) getCached(key string) (*HTMLMeta, bool, error) { + f.cacheMu.Lock() + defer f.cacheMu.Unlock() + + entry, ok := f.cache[key] + if !ok { + return nil, false, nil + } + if !f.now().Before(entry.expiresAt) { + delete(f.cache, key) + return nil, false, nil + } + return cloneHTMLMeta(entry.meta), true, entry.err +} + +func (f *HTMLMetaFetcher) setCached(key string, meta *HTMLMeta, err error, ttl time.Duration) { + f.cacheMu.Lock() + defer f.cacheMu.Unlock() + + now := f.now() + if _, exists := f.cache[key]; !exists && len(f.cache) >= maxCacheEntries { + var oldestKey string + var oldestTime time.Time + for candidateKey, entry := range f.cache { + if oldestKey == "" || entry.storedAt.Before(oldestTime) { + oldestKey = candidateKey + oldestTime = entry.storedAt } } + delete(f.cache, oldestKey) + } + f.cache[key] = cacheEntry{ + meta: cloneHTMLMeta(meta), + err: err, + expiresAt: now.Add(ttl), + storedAt: now, } } + +func normalizeURL(urlStr string) (string, error) { + parsed, err := url.Parse(strings.TrimSpace(urlStr)) + if err != nil { + return "", errors.New("invalid URL format") + } + parsed.Scheme = strings.ToLower(parsed.Scheme) + parsed.Host = strings.ToLower(parsed.Host) + parsed.Fragment = "" + if parsed.Path == "" { + parsed.Path = "/" + } + if (parsed.Scheme == "http" && parsed.Port() == "80") || (parsed.Scheme == "https" && parsed.Port() == "443") { + parsed.Host = parsed.Hostname() + if strings.Contains(parsed.Host, ":") { + parsed.Host = "[" + parsed.Host + "]" + } + } + if err := validateURL(parsed.String()); err != nil { + return "", err + } + return parsed.String(), nil +} + +func finalResponseURL(response *http.Response, fallback string) (*url.URL, error) { + if response.Request != nil && response.Request.URL != nil { + return response.Request.URL, nil + } + parsed, err := url.Parse(fallback) + if err != nil { + return nil, errors.Wrap(err, "invalid response URL") + } + return parsed, nil +} + +func cloneHTMLMeta(meta *HTMLMeta) *HTMLMeta { + if meta == nil { + return nil + } + copy := *meta + return © +} + +// siteImageSource supplies deterministic site-specific thumbnails that outrank +// the low-confidence semantic fallback but lose to metadata the site +// itself declares. +func siteImageSource(pageURL *url.URL) metadataSource { + if pageURL.Hostname() == "www.youtube.com" && pageURL.Path == "/watch" { + if videoID := pageURL.Query().Get("v"); videoID != "" { + return metadataSource{image: fmt.Sprintf("https://img.youtube.com/vi/%s/mqdefault.jpg", url.PathEscape(videoID))} + } + } + return metadataSource{} +} diff --git a/internal/httpgetter/html_meta_extract.go b/internal/httpgetter/html_meta_extract.go new file mode 100644 index 00000000..464d71b9 --- /dev/null +++ b/internal/httpgetter/html_meta_extract.go @@ -0,0 +1,446 @@ +package httpgetter + +import ( + "bytes" + "encoding/json" + "mime" + "net/url" + "strings" + "unicode" + "unicode/utf8" + + "github.com/pkg/errors" + "golang.org/x/net/html" +) + +const ( + maxTitleRunes = 200 + maxDescriptionRunes = 300 + minSemanticTextRunes = 40 + maxImageURLBytes = 2048 +) + +type metadataSource struct { + title string + description string + image string +} + +type documentMetadata struct { + openGraph metadataSource + twitter metadataSource + jsonLD metadataSource + standard metadataSource + semantic metadataSource + baseHref string +} + +func extractDocumentMetadata(document *html.Node) (documentMetadata, string) { + var sources documentMetadata + var oEmbedEndpoint string + var jsonLDScripts []string + + walkNodes(document, func(node *html.Node) { + if node.Type != html.ElementNode { + return + } + switch strings.ToLower(node.Data) { + case "title": + // Foreign-content titles (inline SVG/MathML accessibility labels) + // are not the document title. + if node.Namespace == "" { + setIfEmpty(&sources.standard.title, nodeText(node)) + } + case "base": + setIfEmpty(&sources.baseHref, firstAttribute(node, "href")) + case "meta": + name := strings.ToLower(strings.TrimSpace(firstAttribute(node, "property", "name"))) + content := firstAttribute(node, "content") + switch name { + case "og:title": + setIfEmpty(&sources.openGraph.title, content) + case "og:description": + setIfEmpty(&sources.openGraph.description, content) + case "og:image", "og:image:url", "og:image:secure_url": + setImageIfEmpty(&sources.openGraph.image, content) + case "twitter:title": + setIfEmpty(&sources.twitter.title, content) + case "twitter:description": + setIfEmpty(&sources.twitter.description, content) + case "twitter:image", "twitter:image:src": + setImageIfEmpty(&sources.twitter.image, content) + case "description": + setIfEmpty(&sources.standard.description, content) + case "title", "name", "headline": + setIfEmpty(&sources.standard.title, content) + case "image": + setImageIfEmpty(&sources.standard.image, content) + default: + } + case "link": + rels := strings.Fields(strings.ToLower(firstAttribute(node, "rel"))) + if containsString(rels, "image_src") { + setImageIfEmpty(&sources.standard.image, firstAttribute(node, "href")) + } + if oEmbedEndpoint == "" && containsString(rels, "alternate") { + mediaType := firstAttribute(node, "type") + if isJSONOEmbedMediaType(mediaType) { + oEmbedEndpoint = firstAttribute(node, "href") + } + } + case "script": + if isJSONLDMediaType(firstAttribute(node, "type")) { + jsonLDScripts = append(jsonLDScripts, rawNodeText(node)) + } + default: + } + }) + + sources.jsonLD = extractJSONLDMetadata(jsonLDScripts) + sources.semantic = extractSemanticMetadata(document) + return sources, oEmbedEndpoint +} + +func mergeMetadata(pageURL *url.URL, sources ...metadataSource) *HTMLMeta { + meta := &HTMLMeta{} + for _, source := range sources { + setIfEmpty(&meta.Title, source.title) + setIfEmpty(&meta.Description, source.description) + if meta.Image == "" { + meta.Image = resolveHTTPURL(pageURL, source.image) + } + } + meta.Title = truncateText(normalizeWhitespace(meta.Title), maxTitleRunes) + meta.Description = truncateText(normalizeWhitespace(meta.Description), maxDescriptionRunes) + // A truncated URL would be broken, so oversized image URLs are dropped + // rather than shortened. + if len(meta.Image) > maxImageURLBytes { + meta.Image = "" + } + return meta +} + +func extractOEmbedMetadata(data []byte, baseURL *url.URL) (metadataSource, error) { + var payload struct { + Type string `json:"type"` + Title string `json:"title"` + Description string `json:"description"` + ThumbnailURL string `json:"thumbnail_url"` + URL string `json:"url"` + } + decoder := json.NewDecoder(bytes.NewReader(data)) + if err := decoder.Decode(&payload); err != nil { + return metadataSource{}, errors.Wrap(err, "failed to decode oEmbed response") + } + image := resolveHTTPURL(baseURL, payload.ThumbnailURL) + if image == "" && strings.EqualFold(payload.Type, "photo") { + image = resolveHTTPURL(baseURL, payload.URL) + } + return metadataSource{ + title: payload.Title, + description: payload.Description, + image: image, + }, nil +} + +func extractJSONLDMetadata(scripts []string) metadataSource { + var best metadataSource + bestScore := -1 + for _, script := range scripts { + var value any + decoder := json.NewDecoder(strings.NewReader(script)) + if err := decoder.Decode(&value); err != nil { + continue + } + walkJSONLD(value, func(object map[string]any) { + candidate := metadataSource{ + title: firstJSONString(object["headline"], object["name"]), + description: firstJSONString(object["description"]), + image: jsonLDImage(object["image"]), + } + if candidate.image == "" { + candidate.image = jsonLDImage(object["thumbnailUrl"]) + } + if candidate == (metadataSource{}) { + return + } + score := jsonLDTypeScore(object["@type"]) + if candidate.title != "" { + score++ + } + if candidate.description != "" { + score++ + } + if candidate.image != "" { + score++ + } + if score > bestScore { + best = candidate + bestScore = score + } + }) + } + return best +} + +func walkJSONLD(value any, visit func(map[string]any)) { + switch typed := value.(type) { + case []any: + for _, item := range typed { + walkJSONLD(item, visit) + } + case map[string]any: + visit(typed) + for _, nested := range typed { + walkJSONLD(nested, visit) + } + default: + } +} + +func jsonLDTypeScore(value any) int { + types := jsonLDStrings(value) + for _, value := range types { + switch strings.ToLower(value) { + case "article", "newsarticle", "blogposting", "socialmediaposting", "product", "videoobject", "movie", "book", "event", "recipe": + return 20 + case "creativework", "webpage", "profilepage", "aboutpage", "collectionpage": + return 10 + case "breadcrumblist", "itemlist": + return -10 + } + } + return 0 +} + +func jsonLDImage(value any) string { + switch typed := value.(type) { + case string: + if isPotentialHTTPURL(typed) { + return typed + } + case []any: + for _, item := range typed { + if image := jsonLDImage(item); image != "" { + return image + } + } + case map[string]any: + // "@id" is deliberately excluded: it is a node reference (often the + // page's own URL plus a fragment), not an image location. + for _, candidate := range []any{typed["url"], typed["contentUrl"]} { + if image := jsonLDImage(candidate); image != "" { + return image + } + } + default: + } + return "" +} + +func firstJSONString(values ...any) string { + for _, value := range values { + for _, candidate := range jsonLDStrings(value) { + if strings.TrimSpace(candidate) != "" { + return candidate + } + } + } + return "" +} + +func jsonLDStrings(value any) []string { + switch typed := value.(type) { + case string: + return []string{typed} + case []any: + var result []string + for _, item := range typed { + result = append(result, jsonLDStrings(item)...) + } + return result + } + return nil +} + +func extractSemanticMetadata(document *html.Node) metadataSource { + root := firstElement(document, "main", "article") + if root == nil { + root = firstElement(document, "body") + } + if root == nil { + return metadataSource{} + } + + var source metadataSource + walkNodes(root, func(node *html.Node) { + if node.Type != html.ElementNode { + return + } + switch strings.ToLower(node.Data) { + case "h1": + setIfEmpty(&source.title, nodeText(node)) + case "p": + text := normalizeWhitespace(nodeText(node)) + if source.description == "" && utf8.RuneCountInString(text) >= minSemanticTextRunes { + source.description = text + } + case "img": + setImageIfEmpty(&source.image, firstAttribute(node, "src")) + default: + } + }) + return source +} + +func resolveHTTPURL(baseURL *url.URL, candidate string) string { + candidate = strings.TrimSpace(candidate) + if candidate == "" { + return "" + } + parsed, err := url.Parse(candidate) + if err != nil { + return "" + } + if baseURL != nil { + parsed = baseURL.ResolveReference(parsed) + } + if parsed.Scheme != "http" && parsed.Scheme != "https" { + return "" + } + if parsed.Hostname() == "" { + return "" + } + return parsed.String() +} + +func walkNodes(node *html.Node, visit func(*html.Node)) { + if node == nil { + return + } + visit(node) + for child := node.FirstChild; child != nil; child = child.NextSibling { + walkNodes(child, visit) + } +} + +// firstElement prefers earlier names over document order: all matches for +// names[0] outrank any match for names[1]. +func firstElement(node *html.Node, names ...string) *html.Node { + for _, name := range names { + var result *html.Node + walkNodes(node, func(candidate *html.Node) { + if result != nil || candidate.Type != html.ElementNode { + return + } + if strings.EqualFold(candidate.Data, name) { + result = candidate + } + }) + if result != nil { + return result + } + } + return nil +} + +func firstAttribute(node *html.Node, names ...string) string { + for _, name := range names { + for _, attribute := range node.Attr { + if strings.EqualFold(attribute.Key, name) { + return attribute.Val + } + } + } + return "" +} + +func rawNodeText(node *html.Node) string { + var builder strings.Builder + walkNodes(node, func(candidate *html.Node) { + if candidate.Type == html.TextNode { + builder.WriteString(candidate.Data) + } + }) + return builder.String() +} + +func nodeText(node *html.Node) string { + return normalizeWhitespace(rawNodeText(node)) +} + +func normalizeWhitespace(value string) string { + return strings.Join(strings.FieldsFunc(value, unicode.IsSpace), " ") +} + +func truncateText(value string, maximum int) string { + runes := []rune(value) + if len(runes) <= maximum { + return value + } + return strings.TrimSpace(string(runes[:maximum-1])) + "…" +} + +func setIfEmpty(target *string, value string) { + if *target == "" && strings.TrimSpace(value) != "" { + *target = value + } +} + +func setImageIfEmpty(target *string, value string) { + if *target == "" && isPotentialHTTPURL(value) { + *target = value + } +} + +func isPotentialHTTPURL(value string) bool { + value = strings.TrimSpace(value) + if value == "" { + return false + } + parsed, err := url.Parse(value) + if err != nil { + return false + } + if parsed.Scheme == "" { + return !strings.HasPrefix(value, "//") || parsed.Hostname() != "" + } + return (parsed.Scheme == "http" || parsed.Scheme == "https") && parsed.Hostname() != "" +} + +func containsString(values []string, expected string) bool { + for _, value := range values { + if value == expected { + return true + } + } + return false +} + +func isJSONOEmbedMediaType(value string) bool { + mediaType, _, err := mime.ParseMediaType(value) + return err == nil && strings.EqualFold(mediaType, "application/json+oembed") +} + +func isJSONLDMediaType(value string) bool { + mediaType, _, err := mime.ParseMediaType(value) + return err == nil && strings.EqualFold(mediaType, "application/ld+json") +} + +// resolveDocumentBase applies the document's , falling back to the +// final response URL when the base is missing or not an absolute http(s) URL. +func resolveDocumentBase(pageURL *url.URL, baseHref string) *url.URL { + baseHref = strings.TrimSpace(baseHref) + if baseHref == "" { + return pageURL + } + parsed, err := url.Parse(baseHref) + if err != nil { + return pageURL + } + resolved := pageURL.ResolveReference(parsed) + if (resolved.Scheme != "http" && resolved.Scheme != "https") || resolved.Hostname() == "" { + return pageURL + } + return resolved +} diff --git a/internal/httpgetter/html_meta_test.go b/internal/httpgetter/html_meta_test.go index f10e5ac8..73cda947 100644 --- a/internal/httpgetter/html_meta_test.go +++ b/internal/httpgetter/html_meta_test.go @@ -3,13 +3,20 @@ package httpgetter import ( "context" "errors" + "fmt" "io" "net" "net/http" + "net/url" "strings" + "sync" + "sync/atomic" "testing" + "time" + "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + "golang.org/x/net/html" ) type roundTripFunc func(*http.Request) (*http.Response, error) @@ -18,125 +25,452 @@ func (f roundTripFunc) RoundTrip(req *http.Request) (*http.Response, error) { return f(req) } -func TestGetHTMLMeta(t *testing.T) { - originalHTTPClient := httpClient - t.Cleanup(func() { - httpClient = originalHTTPClient +func TestExtractDocumentMetadataPrecedenceAndFallbacks(t *testing.T) { + tests := []struct { + name string + markup string + expected HTMLMeta + }{ + { + name: "Open Graph wins field by field", + markup: ` + Standard title + + + + + + `, + expected: HTMLMeta{ + Title: "Open Graph title", + Description: "Twitter description", + Image: "https://example.com/twitter.png", + }, + }, + { + name: "JSON-LD graph and image object", + markup: ``, + expected: HTMLMeta{ + Title: "JSON-LD title", + Description: "JSON-LD description", + Image: "https://example.com/posts/images/cover.jpg", + }, + }, + { + name: "nested JSON-LD main entity", + markup: ``, + expected: HTMLMeta{ + Title: "Nested article title", + Description: "Nested article description", + Image: "https://example.com/nested.png", + }, + }, + { + name: "standard metadata", + markup: `Standard title + + + `, + expected: HTMLMeta{ + Title: "Standard title", + Description: "Standard description", + Image: "https://example.com/standard.png", + }, + }, + { + name: "semantic main content", + markup: `
+

Semantic title

+

This paragraph is deliberately long enough to become the semantic preview description.

+ +
`, + expected: HTMLMeta{ + Title: "Semantic title", + Description: "This paragraph is deliberately long enough to become the semantic preview description.", + Image: "https://example.com/posts/semantic.png", + }, + }, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + meta := extractTestMetadata(t, test.markup) + require.Equal(t, test.expected, *meta) + }) + } +} + +func TestExtractJSONLDImageForms(t *testing.T) { + tests := []struct { + name string + image string + want string + }{ + {name: "string", image: `"/one.png"`, want: "https://example.com/one.png"}, + {name: "array", image: `["data:image/png;base64,abc", "/two.png"]`, want: "https://example.com/two.png"}, + {name: "url object", image: `{"url":"/three.png"}`, want: "https://example.com/three.png"}, + {name: "content URL object", image: `{"url":"javascript:alert(1)","contentUrl":"/four.png"}`, want: "https://example.com/four.png"}, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + markup := `` + meta := extractTestMetadata(t, markup) + require.Equal(t, test.want, meta.Image) + }) + } +} + +func TestExtractDocumentMetadataNormalizesAndTruncatesText(t *testing.T) { + markup := ` + ` + meta := extractTestMetadata(t, markup) + + require.Len(t, []rune(meta.Title), maxTitleRunes) + require.NotContains(t, meta.Title, " ") + require.True(t, strings.HasSuffix(meta.Title, "…")) + require.Len(t, []rune(meta.Description), maxDescriptionRunes) + require.True(t, strings.HasSuffix(meta.Description, "…")) +} + +func TestExtractDocumentMetadataSkipsUnsafeImages(t *testing.T) { + markup := ` + + + ` + meta := extractTestMetadata(t, markup) + require.Equal(t, "https://example.com/safe.png", meta.Image) +} + +func TestHTMLMetaFetcherOEmbed(t *testing.T) { + var requests atomic.Int32 + fetcher := newTestFetcher(roundTripFunc(func(req *http.Request) (*http.Response, error) { + requests.Add(1) + switch req.URL.Path { + case "/post": + return htmlResponse(req, ` + + + + + `), nil + case "/oembed": + assert.Contains(t, req.Header.Get("Accept"), "application/json+oembed") + return response(req, http.StatusOK, "application/json+oembed", `{ + "type":"rich","title":"oEmbed title","description":"oEmbed description", + "thumbnail_url":"/oembed.png","html":"" + }`), nil + default: + return nil, errors.New("unexpected request path: " + req.URL.Path) + } + })) + + meta, err := fetcher.Get(context.Background(), "http://93.184.216.34/post") + require.NoError(t, err) + require.Equal(t, &HTMLMeta{ + Title: "oEmbed title", + Description: "oEmbed description", + Image: "http://93.184.216.34/oembed.png", + }, meta) + require.EqualValues(t, 2, requests.Load()) +} + +func TestHTMLMetaFetcherIgnoresInvalidOEmbed(t *testing.T) { + tests := []struct { + name string + endpoint string + contentType string + body string + wantCalls int32 + }{ + {name: "unsupported content type", endpoint: "/oembed", contentType: "text/html", body: `

not JSON

`, wantCalls: 2}, + {name: "oversized response", endpoint: "/oembed", contentType: "application/json", body: strings.Repeat("x", maxOEmbedBytes+1), wantCalls: 2}, + {name: "unsafe endpoint", endpoint: "http://127.0.0.1/oembed", contentType: "application/json", body: `{}`, wantCalls: 1}, + {name: "XML discovery", endpoint: "/oembed", contentType: "application/json", body: `{}`, wantCalls: 1}, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + var calls atomic.Int32 + fetcher := newTestFetcher(roundTripFunc(func(req *http.Request) (*http.Response, error) { + call := calls.Add(1) + if call == 1 { + discoveryType := "application/json+oembed" + if test.name == "XML discovery" { + discoveryType = "text/xml+oembed" + } + return htmlResponse(req, ` + `), nil + } + return response(req, http.StatusOK, test.contentType, test.body), nil + })) + + meta, err := fetcher.Get(context.Background(), "http://93.184.216.34/post") + require.NoError(t, err) + require.Equal(t, "Fallback title", meta.Title) + require.Equal(t, test.wantCalls, calls.Load()) + }) + } +} + +func TestHTMLMetaFetcherBotAwareSPA(t *testing.T) { + transport := roundTripFunc(func(req *http.Request) (*http.Response, error) { + assert.Contains(t, req.Header.Get("Accept"), "text/html") + if req.Header.Get("User-Agent") == defaultLinkPreviewUserAgent { + return htmlResponse(req, ` + `), nil + } + return htmlResponse(req, `Application shell
`), nil }) - httpClient = &http.Client{ - Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) { - require.Equal(t, "http://93.184.216.34/article", req.URL.String()) - return &http.Response{ - StatusCode: http.StatusOK, - Header: http.Header{"Content-Type": []string{"text/html; charset=utf-8"}}, - Body: io.NopCloser(strings.NewReader(` - - - Fallback title - - - - - -ignored -`)), - Request: req, - }, nil - }), - } - - metadata, err := GetHTMLMeta("http://93.184.216.34/article") + fetcher := newTestFetcher(transport) + meta, err := fetcher.Get(context.Background(), "http://93.184.216.34/spa") require.NoError(t, err) - require.Equal(t, HTMLMeta{ - Title: "Open Graph title", - Description: "Open Graph description", - Image: "https://example.com/cover.png", - }, *metadata) + require.Equal(t, "Rich SPA title", meta.Title) + require.Equal(t, "Prerendered for MemosBot", meta.Description) } -func TestGetHTMLMetaWithNameOnly(t *testing.T) { - originalHTTPClient := httpClient - t.Cleanup(func() { - httpClient = originalHTTPClient - }) +func TestHTMLMetaFetcherDefaultHeadersAndDeadline(t *testing.T) { + fetcher := newTestFetcher(roundTripFunc(func(req *http.Request) (*http.Response, error) { + assert.Equal(t, defaultLinkPreviewUserAgent, req.Header.Get("User-Agent")) + deadline, ok := req.Context().Deadline() + assert.True(t, ok) + assert.LessOrEqual(t, time.Until(deadline), fetchTimeout) + return htmlResponse(req, `Title`), nil + })) - httpClient = &http.Client{ - Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) { - require.Equal(t, "http://93.184.216.34/blog", req.URL.String()) - return &http.Response{ - StatusCode: http.StatusOK, - Header: http.Header{"Content-Type": []string{"text/html; charset=utf-8"}}, - Body: io.NopCloser(strings.NewReader(` - - - Sample Page - - -Hello -`)), - Request: req, - }, nil - }), - } - - metadata, err := GetHTMLMeta("http://93.184.216.34/blog") + _, err := fetcher.Get(context.Background(), "http://93.184.216.34/page") require.NoError(t, err) - require.Equal(t, HTMLMeta{ - Title: "Sample Page", - Description: "This description should appear in the link preview.", - Image: "", - }, *metadata) + require.Equal(t, fetchTimeout, newHTTPClient().Timeout) } -func TestGetHTMLMetaWithNameCaseInsensitive(t *testing.T) { - originalHTTPClient := httpClient - t.Cleanup(func() { - httpClient = originalHTTPClient - }) - - httpClient = &http.Client{ - Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) { - require.Equal(t, "http://93.184.216.34/blog", req.URL.String()) - return &http.Response{ - StatusCode: http.StatusOK, - Header: http.Header{"Content-Type": []string{"text/html; charset=utf-8"}}, - Body: io.NopCloser(strings.NewReader(` - - - Sample Page - - -Hello -`)), - Request: req, - }, nil - }), - } - - metadata, err := GetHTMLMeta("http://93.184.216.34/blog") +func TestSiteImageSourceEscapesYouTubeVideoID(t *testing.T) { + pageURL, err := url.Parse("https://www.youtube.com/watch?v=video%2Fsegment%3Fvariant") require.NoError(t, err) - require.Equal(t, HTMLMeta{ - Title: "Sample Page", - Description: "Case insensitive description match.", - Image: "", - }, *metadata) + + source := siteImageSource(pageURL) + require.Equal(t, "https://img.youtube.com/vi/video%2Fsegment%3Fvariant/mqdefault.jpg", source.image) } -func TestGetHTMLMetaForInternal(t *testing.T) { - // test for internal IP - if _, err := GetHTMLMeta("http://192.168.0.1"); !errors.Is(err, ErrInternalIP) { - t.Errorf("Expected error for internal IP, got %v", err) +func TestHTMLMetaFetcherResolvesAgainstFinalResponseURL(t *testing.T) { + fetcher := newTestFetcher(roundTripFunc(func(req *http.Request) (*http.Response, error) { + finalRequest := req.Clone(req.Context()) + finalRequest.URL, _ = url.Parse("https://cdn.example.com/final/path/") + return htmlResponse(finalRequest, ``), nil + })) + + meta, err := fetcher.Get(context.Background(), "http://93.184.216.34/original") + require.NoError(t, err) + require.Equal(t, "https://cdn.example.com/final/path/cover.png", meta.Image) +} + +func TestHTMLMetaFetcherValidatesResponse(t *testing.T) { + tests := []struct { + name string + status int + contentType string + }{ + {name: "non-2xx", status: http.StatusNotFound, contentType: "text/html"}, + {name: "unsupported content type", status: http.StatusOK, contentType: "application/pdf"}, } - // test for resolved internal IP - if _, err := GetHTMLMeta("http://localhost"); !errors.Is(err, ErrInternalIP) { - t.Errorf("Expected error for resolved internal IP, got %v", err) + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + fetcher := newTestFetcher(roundTripFunc(func(req *http.Request) (*http.Response, error) { + return response(req, test.status, test.contentType, "body"), nil + })) + _, err := fetcher.Get(context.Background(), "http://93.184.216.34/page") + require.Error(t, err) + }) } } -func TestHTTPClientHasTimeout(t *testing.T) { - require.NotZero(t, httpClient.Timeout) +func TestHTMLMetaFetcherBoundsHTMLResponse(t *testing.T) { + body := `Bounded title` + strings.Repeat("x", maxHTMLMetaBytes) + `` + reader := &countingReadCloser{Reader: strings.NewReader(body)} + fetcher := newTestFetcher(roundTripFunc(func(req *http.Request) (*http.Response, error) { + result := response(req, http.StatusOK, "text/html", "") + result.Body = reader + return result, nil + })) + + meta, err := fetcher.Get(context.Background(), "http://93.184.216.34/page") + require.NoError(t, err) + require.Equal(t, "Bounded title", meta.Title) + require.LessOrEqual(t, reader.bytesRead, int64(maxHTMLMetaBytes)) +} + +func TestHTMLMetaFetcherContextCancellation(t *testing.T) { + requestStarted := make(chan struct{}) + fetcher := newTestFetcher(roundTripFunc(func(req *http.Request) (*http.Response, error) { + close(requestStarted) + <-req.Context().Done() + return nil, req.Context().Err() + })) + + ctx, cancel := context.WithCancel(context.Background()) + result := make(chan error, 1) + go func() { + _, err := fetcher.Get(ctx, "http://93.184.216.34/page") + result <- err + }() + <-requestStarted + cancel() + require.ErrorIs(t, <-result, context.Canceled) +} + +func TestHTMLMetaFetcherCache(t *testing.T) { + var now = time.Date(2026, time.August, 17, 0, 0, 0, 0, time.UTC) + var successCalls atomic.Int32 + successFetcher := newTestFetcher(roundTripFunc(func(req *http.Request) (*http.Response, error) { + successCalls.Add(1) + return htmlResponse(req, `Cached title`), nil + })) + successFetcher.now = func() time.Time { return now } + + first, err := successFetcher.Get(context.Background(), "HTTP://93.184.216.34:80/page#one") + require.NoError(t, err) + first.Title = "mutated" + second, err := successFetcher.Get(context.Background(), "http://93.184.216.34/page#two") + require.NoError(t, err) + require.Equal(t, "Cached title", second.Title) + require.EqualValues(t, 1, successCalls.Load()) + + now = now.Add(successCacheTTL + time.Second) + _, err = successFetcher.Get(context.Background(), "http://93.184.216.34/page") + require.NoError(t, err) + require.EqualValues(t, 2, successCalls.Load()) + + var failureCalls atomic.Int32 + failureFetcher := newTestFetcher(roundTripFunc(func(req *http.Request) (*http.Response, error) { + failureCalls.Add(1) + return response(req, http.StatusBadGateway, "text/html", "failure"), nil + })) + failureFetcher.now = func() time.Time { return now } + for range 2 { + _, err := failureFetcher.Get(context.Background(), "http://93.184.216.34/failure") + require.Error(t, err) + } + require.EqualValues(t, 1, failureCalls.Load()) + now = now.Add(failureCacheTTL + time.Second) + _, err = failureFetcher.Get(context.Background(), "http://93.184.216.34/failure") + require.Error(t, err) + require.EqualValues(t, 2, failureCalls.Load()) +} + +func TestHTMLMetaFetcherBoundsCache(t *testing.T) { + fetcher := NewHTMLMetaFetcher() + now := time.Date(2026, time.August, 17, 0, 0, 0, 0, time.UTC) + fetcher.now = func() time.Time { return now } + for index := range maxCacheEntries + 1 { + fetcher.setCached(fmt.Sprintf("https://example.com/%d", index), &HTMLMeta{Title: "title"}, nil, successCacheTTL) + now = now.Add(time.Second) + } + + require.Len(t, fetcher.cache, maxCacheEntries) + _, exists := fetcher.cache["https://example.com/0"] + require.False(t, exists) +} + +func TestHTMLMetaFetcherCoalescesRequests(t *testing.T) { + var calls atomic.Int32 + started := make(chan struct{}) + release := make(chan struct{}) + fetcher := newTestFetcher(roundTripFunc(func(req *http.Request) (*http.Response, error) { + if calls.Add(1) == 1 { + close(started) + } + <-release + return htmlResponse(req, `Shared title`), nil + })) + + const requestCount = 20 + var waitGroup sync.WaitGroup + errorsChannel := make(chan error, requestCount) + for range requestCount { + waitGroup.Add(1) + go func() { + defer waitGroup.Done() + meta, err := fetcher.Get(context.Background(), "http://93.184.216.34/shared") + if err == nil && meta.Title != "Shared title" { + err = errors.New("unexpected metadata") + } + errorsChannel <- err + }() + } + <-started + close(release) + waitGroup.Wait() + close(errorsChannel) + for err := range errorsChannel { + require.NoError(t, err) + } + require.EqualValues(t, 1, calls.Load()) +} + +func TestHTMLMetaFetcherLimitsConcurrency(t *testing.T) { + var active atomic.Int32 + var maximum atomic.Int32 + entered := make(chan struct{}, 12) + release := make(chan struct{}) + fetcher := newTestFetcher(roundTripFunc(func(req *http.Request) (*http.Response, error) { + current := active.Add(1) + for { + previous := maximum.Load() + if current <= previous || maximum.CompareAndSwap(previous, current) { + break + } + } + entered <- struct{}{} + <-release + active.Add(-1) + return htmlResponse(req, `Title`), nil + })) + + const requestCount = 12 + var waitGroup sync.WaitGroup + for index := range requestCount { + waitGroup.Add(1) + go func() { + defer waitGroup.Done() + _, _ = fetcher.Get(context.Background(), "http://93.184.216.34/page?id="+string(rune('a'+index))) + }() + } + for range maxConcurrentFetches { + <-entered + } + select { + case <-entered: + t.Fatal("more than eight uncached fetches started concurrently") + case <-time.After(100 * time.Millisecond): + } + close(release) + waitGroup.Wait() + require.EqualValues(t, maxConcurrentFetches, maximum.Load()) +} + +func TestHTMLMetaFetcherRejectsUnsafeURLsAndRedirects(t *testing.T) { + fetcher := NewHTMLMetaFetcher() + _, err := fetcher.Get(context.Background(), "http://192.168.0.1/page") + require.ErrorIs(t, err, ErrInternalIP) + + redirect, err := http.NewRequest(http.MethodGet, "http://127.0.0.1/private", nil) + require.NoError(t, err) + err = newHTTPClient().CheckRedirect(redirect, nil) + require.ErrorIs(t, err, ErrInternalIP) } func TestSecureDialContextRejectsResolvedInternalIP(t *testing.T) { @@ -187,3 +521,48 @@ func TestSecureDialContextDialsResolvedIP(t *testing.T) { require.NotNil(t, conn) require.Equal(t, "93.184.216.34:80", dialedAddress) } + +func extractTestMetadata(t *testing.T, markup string) *HTMLMeta { + t.Helper() + document, err := html.Parse(strings.NewReader(markup)) + require.NoError(t, err) + sources, _ := extractDocumentMetadata(document) + pageURL, err := url.Parse("https://example.com/posts/page") + require.NoError(t, err) + return mergeMetadata(pageURL, metadataSource{}, sources.openGraph, sources.twitter, sources.jsonLD, sources.standard, sources.semantic) +} + +func newTestFetcher(transport http.RoundTripper) *HTMLMetaFetcher { + fetcher := NewHTMLMetaFetcher() + fetcher.client = &http.Client{Transport: transport} + return fetcher +} + +func htmlResponse(request *http.Request, body string) *http.Response { + return response(request, http.StatusOK, "text/html; charset=utf-8", body) +} + +func response(request *http.Request, statusCode int, contentType, body string) *http.Response { + return &http.Response{ + StatusCode: statusCode, + Status: http.StatusText(statusCode), + Header: http.Header{"Content-Type": []string{contentType}}, + Body: io.NopCloser(strings.NewReader(body)), + Request: request, + } +} + +type countingReadCloser struct { + *strings.Reader + bytesRead int64 +} + +func (reader *countingReadCloser) Read(buffer []byte) (int, error) { + count, err := reader.Reader.Read(buffer) + reader.bytesRead += int64(count) + return count, err +} + +func (*countingReadCloser) Close() error { + return nil +} diff --git a/server/router/api/v1/link_metadata_test.go b/server/router/api/v1/link_metadata_test.go index c9c2ce84..b3873333 100644 --- a/server/router/api/v1/link_metadata_test.go +++ b/server/router/api/v1/link_metadata_test.go @@ -5,6 +5,7 @@ import ( "fmt" "testing" + "github.com/pkg/errors" "github.com/stretchr/testify/require" "google.golang.org/grpc/codes" "google.golang.org/grpc/status" @@ -13,22 +14,29 @@ import ( v1pb "github.com/usememos/memos/proto/gen/api/v1" ) -func TestGetLinkMetadata(t *testing.T) { - originalFetchHTMLMeta := fetchHTMLMeta - t.Cleanup(func() { - fetchHTMLMeta = originalFetchHTMLMeta - }) +type fakeLinkMetadataFetcher func(context.Context, string) (*httpgetter.HTMLMeta, error) - fetchHTMLMeta = func(url string) (*httpgetter.HTMLMeta, error) { - require.Equal(t, "https://example.com/article", url) - return &httpgetter.HTMLMeta{ - Title: "Example title", - Description: "Example description", - Image: "https://example.com/cover.png", - }, nil +type linkMetadataContextKey struct{} + +func (fetch fakeLinkMetadataFetcher) Get(ctx context.Context, url string) (*httpgetter.HTMLMeta, error) { + return fetch(ctx, url) +} + +func TestGetLinkMetadata(t *testing.T) { + requestContext := context.WithValue(context.Background(), linkMetadataContextKey{}, "context value") + service := &APIV1Service{ + linkMetadataFetcher: fakeLinkMetadataFetcher(func(ctx context.Context, url string) (*httpgetter.HTMLMeta, error) { + require.Same(t, requestContext, ctx) + require.Equal(t, "https://example.com/article", url) + return &httpgetter.HTMLMeta{ + Title: "Example title", + Description: "Example description", + Image: "https://example.com/cover.png", + }, nil + }), } - metadata, err := (&APIV1Service{}).GetLinkMetadata(context.Background(), &v1pb.GetLinkMetadataRequest{ + metadata, err := service.GetLinkMetadata(requestContext, &v1pb.GetLinkMetadataRequest{ Url: "https://example.com/article", }) require.NoError(t, err) @@ -44,31 +52,41 @@ func TestGetLinkMetadataEmptyURL(t *testing.T) { require.Equal(t, codes.InvalidArgument, status.Code(err)) } +func TestGetLinkMetadataFetchError(t *testing.T) { + service := &APIV1Service{ + linkMetadataFetcher: fakeLinkMetadataFetcher(func(context.Context, string) (*httpgetter.HTMLMeta, error) { + return nil, errors.New("fetch failed") + }), + } + _, err := service.GetLinkMetadata(context.Background(), &v1pb.GetLinkMetadataRequest{Url: "https://example.com"}) + require.Error(t, err) + require.Equal(t, codes.InvalidArgument, status.Code(err)) + require.Contains(t, err.Error(), "failed to fetch link metadata") +} + func TestGetLinkMetadataInternalURL(t *testing.T) { - _, err := (&APIV1Service{}).GetLinkMetadata(context.Background(), &v1pb.GetLinkMetadataRequest{ + service := &APIV1Service{linkMetadataFetcher: httpgetter.NewHTMLMetaFetcher()} + _, err := service.GetLinkMetadata(context.Background(), &v1pb.GetLinkMetadataRequest{ Url: "http://192.168.0.1", }) require.Error(t, err) require.Equal(t, codes.InvalidArgument, status.Code(err)) } -func TestBatchGetLinkMetadata(t *testing.T) { - originalFetchHTMLMeta := fetchHTMLMeta - t.Cleanup(func() { - fetchHTMLMeta = originalFetchHTMLMeta - }) - +func TestBatchGetLinkMetadataPreservesOrder(t *testing.T) { var fetchedURLs []string - fetchHTMLMeta = func(url string) (*httpgetter.HTMLMeta, error) { - fetchedURLs = append(fetchedURLs, url) - return &httpgetter.HTMLMeta{ - Title: fmt.Sprintf("Title for %s", url), - Description: fmt.Sprintf("Description for %s", url), - Image: fmt.Sprintf("%s/cover.png", url), - }, nil + service := &APIV1Service{ + linkMetadataFetcher: fakeLinkMetadataFetcher(func(_ context.Context, url string) (*httpgetter.HTMLMeta, error) { + fetchedURLs = append(fetchedURLs, url) + return &httpgetter.HTMLMeta{ + Title: fmt.Sprintf("Title for %s", url), + Description: fmt.Sprintf("Description for %s", url), + Image: fmt.Sprintf("%s/cover.png", url), + }, nil + }), } - response, err := (&APIV1Service{}).BatchGetLinkMetadata(context.Background(), &v1pb.BatchGetLinkMetadataRequest{ + response, err := service.BatchGetLinkMetadata(context.Background(), &v1pb.BatchGetLinkMetadataRequest{ Urls: []string{ "https://example.com/one", "https://example.com/two", @@ -83,6 +101,26 @@ func TestBatchGetLinkMetadata(t *testing.T) { require.Equal(t, "Title for https://example.com/two", response.LinkMetadata[1].Title) } +func TestBatchGetLinkMetadataStopsOnError(t *testing.T) { + var fetchedURLs []string + service := &APIV1Service{ + linkMetadataFetcher: fakeLinkMetadataFetcher(func(_ context.Context, url string) (*httpgetter.HTMLMeta, error) { + fetchedURLs = append(fetchedURLs, url) + if url == "https://example.com/two" { + return nil, errors.New("fetch failed") + } + return &httpgetter.HTMLMeta{Title: url}, nil + }), + } + + _, err := service.BatchGetLinkMetadata(context.Background(), &v1pb.BatchGetLinkMetadataRequest{ + Urls: []string{"https://example.com/one", "https://example.com/two", "https://example.com/three"}, + }) + require.Error(t, err) + require.Equal(t, codes.InvalidArgument, status.Code(err)) + require.Equal(t, []string{"https://example.com/one", "https://example.com/two"}, fetchedURLs) +} + func TestBatchGetLinkMetadataEmptyURLs(t *testing.T) { _, err := (&APIV1Service{}).BatchGetLinkMetadata(context.Background(), &v1pb.BatchGetLinkMetadataRequest{}) require.Error(t, err) diff --git a/server/router/api/v1/memo_service.go b/server/router/api/v1/memo_service.go index 1bd27883..9179b1cd 100644 --- a/server/router/api/v1/memo_service.go +++ b/server/router/api/v1/memo_service.go @@ -14,7 +14,6 @@ import ( "google.golang.org/protobuf/proto" "google.golang.org/protobuf/types/known/emptypb" - "github.com/usememos/memos/internal/httpgetter" v1pb "github.com/usememos/memos/proto/gen/api/v1" storepb "github.com/usememos/memos/proto/gen/store" "github.com/usememos/memos/server/access" @@ -28,8 +27,6 @@ type suppressSSEKey struct{} const maxBatchGetLinkMetadata = 10 -var fetchHTMLMeta = httpgetter.GetHTMLMeta - func withSuppressSSE(ctx context.Context) context.Context { return context.WithValue(ctx, suppressSSEKey{}, true) } diff --git a/server/router/api/v1/memo_service_link_metadata.go b/server/router/api/v1/memo_service_link_metadata.go index aa77a80c..d641355c 100644 --- a/server/router/api/v1/memo_service_link_metadata.go +++ b/server/router/api/v1/memo_service_link_metadata.go @@ -7,16 +7,21 @@ import ( "google.golang.org/grpc/codes" "google.golang.org/grpc/status" + "github.com/usememos/memos/internal/httpgetter" v1pb "github.com/usememos/memos/proto/gen/api/v1" ) +type linkMetadataFetcher interface { + Get(context.Context, string) (*httpgetter.HTMLMeta, error) +} + // GetLinkMetadata gets metadata for a link. -func (*APIV1Service) GetLinkMetadata(_ context.Context, request *v1pb.GetLinkMetadataRequest) (*v1pb.LinkMetadata, error) { - return getLinkMetadata(request.GetUrl()) +func (s *APIV1Service) GetLinkMetadata(ctx context.Context, request *v1pb.GetLinkMetadataRequest) (*v1pb.LinkMetadata, error) { + return s.buildLinkMetadata(ctx, request.GetUrl()) } // BatchGetLinkMetadata gets metadata for links. -func (*APIV1Service) BatchGetLinkMetadata(_ context.Context, request *v1pb.BatchGetLinkMetadataRequest) (*v1pb.BatchGetLinkMetadataResponse, error) { +func (s *APIV1Service) BatchGetLinkMetadata(ctx context.Context, request *v1pb.BatchGetLinkMetadataRequest) (*v1pb.BatchGetLinkMetadataResponse, error) { if len(request.Urls) == 0 { return nil, status.Errorf(codes.InvalidArgument, "urls are required") } @@ -26,7 +31,7 @@ func (*APIV1Service) BatchGetLinkMetadata(_ context.Context, request *v1pb.Batch linkMetadata := make([]*v1pb.LinkMetadata, 0, len(request.Urls)) for _, url := range request.Urls { - metadata, err := getLinkMetadata(url) + metadata, err := s.buildLinkMetadata(ctx, url) if err != nil { return nil, err } @@ -38,12 +43,12 @@ func (*APIV1Service) BatchGetLinkMetadata(_ context.Context, request *v1pb.Batch }, nil } -func getLinkMetadata(inputURL string) (*v1pb.LinkMetadata, error) { +func (s *APIV1Service) buildLinkMetadata(ctx context.Context, inputURL string) (*v1pb.LinkMetadata, error) { url := strings.TrimSpace(inputURL) if url == "" { return nil, status.Errorf(codes.InvalidArgument, "url is required") } - htmlMeta, err := fetchHTMLMeta(url) + htmlMeta, err := s.linkMetadataFetcher.Get(ctx, url) if err != nil { return nil, status.Errorf(codes.InvalidArgument, "failed to fetch link metadata: %v", err) } diff --git a/server/router/api/v1/v1.go b/server/router/api/v1/v1.go index b9cba332..3e878f89 100644 --- a/server/router/api/v1/v1.go +++ b/server/router/api/v1/v1.go @@ -11,6 +11,7 @@ import ( "golang.org/x/sync/semaphore" "google.golang.org/protobuf/encoding/protojson" + "github.com/usememos/memos/internal/httpgetter" "github.com/usememos/memos/internal/markdown" "github.com/usememos/memos/internal/profile" v1pb "github.com/usememos/memos/proto/gen/api/v1" @@ -47,14 +48,17 @@ type APIV1Service struct { // instanceStatsCache memoizes GetInstanceStats results for instanceStatsCacheTTL. instanceStatsCache instanceStatsCache + + linkMetadataFetcher linkMetadataFetcher } +// NewAPIV1Service creates an API v1 service with its shared dependencies. func NewAPIV1Service(secret string, profile *profile.Profile, store *store.Store) *APIV1Service { markdownService := markdown.NewService( markdown.WithTagExtension(), markdown.WithMentionExtension(), ) - return &APIV1Service{ + service := &APIV1Service{ Secret: secret, Profile: profile, Store: store, @@ -64,6 +68,8 @@ func NewAPIV1Service(secret string, profile *profile.Profile, store *store.Store thumbnailSemaphore: semaphore.NewWeighted(3), // Limit to 3 concurrent thumbnail generations imageProcessingSemaphore: semaphore.NewWeighted(2), } + service.linkMetadataFetcher = httpgetter.NewHTMLMetaFetcher() + return service } // newGatewayMarshaler mirrors grpc-gateway's default JSON marshaler with one