Files
grendervill a9b1073877 feat(unfurl): превью ссылок с защитой от SSRF и кэшем в БД (Фаза 7)
Сервер сам загружает заголовок, описание и картинку страницы по ссылке из
сообщения и отдаёт клиенту готовую карточку.

Безопасность (главное здесь):
- только http/https и без userinfo; запрет петли, частных сетей, link-local
  (169.254.169.254), CGNAT, multicast и IPv4-mapped вариантов;
- проверка идёт по адресу, к которому реально открывается TCP
  (`net.Dialer.Control`), поэтому подмена DNS между проверкой и соединением
  (DNS rebinding) ничего не даёт;
- не больше 3 редиректов, каждый хоп проверяется заново; таймаут 5 с, тело
  ≤ 512 КБ, только `text/html`; прокси из окружения игнорируются, cookie и
  авторизация не отправляются; в логи попадают только хост и код причины;
- картинка по ссылке не скачивается — проверяется лишь её URL: экономия CPU на
  1 vCPU и минус класс атак через декодирование.

Кэш: таблица `link_previews` (миграция 00022, ключ — sha256 нормализованного
URL), TTL по статусу (ok — сутки, empty/blocked — час, error — 10 минут).
Ручка `GET /api/v1/link-previews?url=…` отвечает статусом
(ok/empty/blocked/error) и карточкой только при ok; 20 новых загрузок в минуту
на пользователя, кэшированные ответы лимит не тратят. `UNFURL_ENABLED=false`
выключает функцию целиком, `features.unfurl_enabled` виден в `/meta`. Retention
убирает истёкшие записи кэша.

Тесты: 21 в `internal/unfurl` (включая DNS rebinding через локальный
DNS-сервер, редирект во внутреннюю сеть, таймаут, лимиты размера и типа),
ручки, store и миграция.
2026-09-26 16:15:00 +03:00

665 lines
24 KiB
Go
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
package unfurl
import (
"bytes"
"context"
"errors"
"io"
"log/slog"
"net"
"net/http"
"net/http/httptest"
"net/netip"
"os"
"strconv"
"strings"
"sync/atomic"
"testing"
"time"
"unicode/utf8"
"golang.org/x/net/dns/dnsmessage"
)
// testFetcher собирает загрузчик для httptest-серверов: они живут на петле,
// поэтому без AllowPrivate адрес 127.0.0.1 был бы запрещён.
func testFetcher(t *testing.T, mutate func(*Options)) *Fetcher {
t.Helper()
options := Options{Timeout: 2 * time.Second, AllowPrivate: true}
if mutate != nil {
mutate(&options)
}
return NewFetcher(options)
}
// serveHTML поднимает сервер, отдающий заданный Content-Type и тело.
func serveHTML(t *testing.T, contentType string, body string) *httptest.Server {
t.Helper()
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
if contentType != "" {
w.Header().Set("Content-Type", contentType)
}
_, _ = io.WriteString(w, body)
}))
t.Cleanup(server.Close)
return server
}
func TestNormalize(t *testing.T) {
cases := []struct {
name string
raw string
want string
ok bool
}{
{"регистр хоста и фрагмент", "https://Example.COM/Path?q=1#frag", "https://example.com/Path?q=1", true},
{"пустой путь", "http://example.com", "http://example.com/", true},
{"порт сохраняется", "http://example.com:8080/a", "http://example.com:8080/a", true},
{"пробелы по краям", " https://example.com/a ", "https://example.com/a", true},
{"ipv6-литерал", "http://[::1]:8080/x", "http://[::1]:8080/x", true},
{"file", "file:///etc/passwd", "", false},
{"javascript", "javascript:alert(1)", "", false},
{"data", "data:text/html,<h1>x</h1>", "", false},
{"ftp", "ftp://example.com/x", "", false},
{"gopher", "gopher://example.com/", "", false},
{"userinfo", "http://user:pass@example.com/", "", false},
{"пустой хост", "http://", "", false},
{"пустая строка", "", "", false},
{"относительная ссылка", "/only/path", "", false},
}
for _, testCase := range cases {
t.Run(testCase.name, func(t *testing.T) {
got, ok := Normalize(testCase.raw)
if ok != testCase.ok {
t.Fatalf("Normalize(%q) ok = %v, want %v", testCase.raw, ok, testCase.ok)
}
if got != testCase.want {
t.Errorf("Normalize(%q) = %q, want %q", testCase.raw, got, testCase.want)
}
})
}
}
func TestURLHash(t *testing.T) {
first := URLHash("https://example.com/a")
second := URLHash("https://example.com/a")
other := URLHash("https://example.com/b")
if len(first) != 64 {
t.Errorf("длина хэша = %d, want 64", len(first))
}
if strings.Trim(first, "0123456789abcdef") != "" {
t.Errorf("хэш %q не является hex-строкой", first)
}
if first != second {
t.Errorf("хэш не детерминирован: %q != %q", first, second)
}
if first == other {
t.Errorf("разные URL дали одинаковый хэш %q", first)
}
}
func TestFetchOpenGraph(t *testing.T) {
const page = `<!doctype html>
<html><head>
<meta property="og:title" content="Заголовок &amp; сущность">
<meta property="OG:DESCRIPTION" content="Описание страницы">
<meta property="og:site_name" content="Пример">
<meta property="og:image" content="/img/pic.png">
</head><body>тело</body></html>`
var sawRequest atomic.Bool
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if got := r.Header.Get("User-Agent"); !strings.HasPrefix(got, defaultUserAgent) {
t.Errorf("User-Agent = %q, want префикс %q", got, defaultUserAgent)
}
if got := r.Header.Get("Accept"); got != acceptHTML {
t.Errorf("Accept = %q, want %q", got, acceptHTML)
}
if got := r.Header.Get("Cookie"); got != "" {
t.Errorf("запрос ушёл с Cookie: %q", got)
}
if got := r.Header.Get("Authorization"); got != "" {
t.Errorf("запрос ушёл с Authorization: %q", got)
}
sawRequest.Store(true)
w.Header().Set("Content-Type", "text/html; charset=utf-8")
_, _ = io.WriteString(w, page)
}))
t.Cleanup(server.Close)
preview, err := testFetcher(t, nil).Fetch(context.Background(), server.URL+"/page")
if err != nil {
t.Fatalf("Fetch() вернул ошибку: %v", err)
}
if !sawRequest.Load() {
t.Fatal("сервер не получил запрос")
}
if preview.URL != server.URL+"/page" {
t.Errorf("URL = %q, want %q", preview.URL, server.URL+"/page")
}
if preview.Title != "Заголовок & сущность" {
t.Errorf("Title = %q", preview.Title)
}
if preview.Description != "Описание страницы" {
t.Errorf("Description = %q", preview.Description)
}
if preview.SiteName != "Пример" {
t.Errorf("SiteName = %q", preview.SiteName)
}
if want := server.URL + "/img/pic.png"; preview.ImageURL != want {
t.Errorf("ImageURL = %q, want %q", preview.ImageURL, want)
}
}
func TestFetchFallbackToTitleAndDescription(t *testing.T) {
const page = `<html><head>
<title>
Заголовок
в несколько строк
</title>
<meta name="description" content=" Описание
страницы ">
</head><body></body></html>`
server := serveHTML(t, "text/html", page)
preview, err := testFetcher(t, nil).Fetch(context.Background(), server.URL)
if err != nil {
t.Fatalf("Fetch() вернул ошибку: %v", err)
}
if preview.Title != "Заголовок в несколько строк" {
t.Errorf("Title = %q", preview.Title)
}
if preview.Description != "Описание страницы" {
t.Errorf("Description = %q", preview.Description)
}
}
func TestFetchTwitterFallback(t *testing.T) {
const page = `<html><head>
<meta name="twitter:title" content="Твиттер-заголовок">
<meta name="twitter:description" content="Твиттер-описание">
<meta name="twitter:image" content="img/tw.png">
</head></html>`
server := serveHTML(t, "text/html", page)
preview, err := testFetcher(t, nil).Fetch(context.Background(), server.URL+"/dir/page")
if err != nil {
t.Fatalf("Fetch() вернул ошибку: %v", err)
}
if preview.Title != "Твиттер-заголовок" {
t.Errorf("Title = %q", preview.Title)
}
if preview.Description != "Твиттер-описание" {
t.Errorf("Description = %q", preview.Description)
}
// Относительная картинка разрешается от каталога страницы.
if want := server.URL + "/dir/img/tw.png"; preview.ImageURL != want {
t.Errorf("ImageURL = %q, want %q", preview.ImageURL, want)
}
}
func TestFetchRejectsNonHTML(t *testing.T) {
cases := []struct {
name string
contentType string
}{
{"png", "image/png"},
{"pdf", "application/pdf"},
{"json", "application/json; charset=utf-8"},
{"пустой", ""},
}
for _, testCase := range cases {
t.Run(testCase.name, func(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
// Пустое значение всё равно выставляем в карту заголовков:
// так сервер не подставляет Content-Type по содержимому.
w.Header()["Content-Type"] = []string{testCase.contentType}
_, _ = io.WriteString(w, `<html><head><title>нет</title></head></html>`)
}))
t.Cleanup(server.Close)
_, err := testFetcher(t, nil).Fetch(context.Background(), server.URL)
if !errors.Is(err, ErrContentType) {
t.Fatalf("err = %v, want ErrContentType", err)
}
})
}
}
func TestFetchRejectsTooLarge(t *testing.T) {
body := `<html><head><title>большая страница</title></head><body>` +
strings.Repeat("x", 4096) + `</body></html>`
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.Header().Set("Content-Type", "text/html")
// Content-Length задаём явно: по нему размер проверяется до чтения тела.
w.Header().Set("Content-Length", strconv.Itoa(len(body)))
_, _ = io.WriteString(w, body)
}))
t.Cleanup(server.Close)
fetcher := testFetcher(t, func(options *Options) { options.MaxBytes = 64 })
_, err := fetcher.Fetch(context.Background(), server.URL)
if !errors.Is(err, ErrTooLarge) {
t.Fatalf("err = %v, want ErrTooLarge", err)
}
}
// TestFetchTruncatedBodyKeepsMetadata: без Content-Length тело читается до
// MaxBytes, и уже разобранные метаданные не теряются.
func TestFetchTruncatedBodyKeepsMetadata(t *testing.T) {
head := `<html><head><title>Короткий заголовок</title></head><body>`
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.Header().Set("Content-Type", "text/html")
_, _ = io.WriteString(w, head)
w.(http.Flusher).Flush() // без Content-Length: ответ пойдёт chunked
_, _ = io.WriteString(w, strings.Repeat("x", 8192))
}))
t.Cleanup(server.Close)
fetcher := testFetcher(t, func(options *Options) { options.MaxBytes = 96 })
preview, err := fetcher.Fetch(context.Background(), server.URL)
if err != nil {
t.Fatalf("Fetch() вернул ошибку: %v", err)
}
if preview.Title != "Короткий заголовок" {
t.Errorf("Title = %q, want %q", preview.Title, "Короткий заголовок")
}
}
func TestFetchStopsAtHead(t *testing.T) {
const page = `<html><head></head><body><title>Из тела</title>
<meta property="og:title" content="Тоже из тела"></body></html>`
server := serveHTML(t, "text/html", page)
preview, err := testFetcher(t, nil).Fetch(context.Background(), server.URL)
if err != nil {
t.Fatalf("Fetch() вернул ошибку: %v", err)
}
if !preview.Empty() {
t.Errorf("превью = %+v, want пустое: тело страницы не разбираем", preview)
}
if preview.URL != server.URL+"/" {
t.Errorf("URL = %q, want %q", preview.URL, server.URL+"/")
}
}
// TestFetchBlocksInternalAddresses: запросы на внутренние адреса отбиваются до
// соединения, серверы поднимать не нужно.
func TestFetchBlocksInternalAddresses(t *testing.T) {
urls := []string{
"http://127.0.0.1:1/",
"http://127.255.255.254/",
"http://169.254.169.254/latest/meta-data/",
"http://10.1.2.3/",
"http://192.168.1.1/",
"http://172.16.0.1/",
"http://172.31.255.254/",
"http://0.0.0.0/",
"http://100.64.0.1/",
"http://[::1]:8080/",
"http://[::ffff:10.0.0.1]/",
"http://[fe80::1]/",
"http://[::]/",
}
fetcher := NewFetcher(Options{Timeout: time.Second})
start := time.Now()
for _, raw := range urls {
t.Run(raw, func(t *testing.T) {
_, err := fetcher.Fetch(context.Background(), raw)
if !errors.Is(err, ErrBlocked) {
t.Fatalf("Fetch(%q) err = %v, want ErrBlocked", raw, err)
}
})
}
// Проверка идёт до соединения: весь набор укладывается в доли секунды.
if elapsed := time.Since(start); elapsed > 2*time.Second {
t.Errorf("проверка адресов заняла %v — похоже, запросы всё-таки уходили в сеть", elapsed)
}
}
func TestFetchAllowPrivateDisabledBlocksLoopback(t *testing.T) {
server := serveHTML(t, "text/html", `<html><head><title>привет</title></head></html>`)
fetcher := NewFetcher(Options{Timeout: time.Second})
_, err := fetcher.Fetch(context.Background(), server.URL)
if !errors.Is(err, ErrBlocked) {
t.Fatalf("err = %v, want ErrBlocked", err)
}
}
func TestFetchRejectsUnsupportedURL(t *testing.T) {
fetcher := NewFetcher(Options{Timeout: time.Second})
for _, raw := range []string{"file:///etc/passwd", "javascript:alert(1)", "http://user:pass@example.com/"} {
if _, err := fetcher.Fetch(context.Background(), raw); !errors.Is(err, ErrUnsupported) {
t.Errorf("Fetch(%q) err = %v, want ErrUnsupported", raw, err)
}
}
}
func TestFetchRedirectToBlockedAddress(t *testing.T) {
// Петля в тестовом режиме разрешена (на ней стоят httptest-серверы),
// поэтому цель перехода — заведомо запрещённый частный адрес.
const blocked = "http://10.1.2.3:1/"
var mux http.ServeMux
mux.HandleFunc("/one", func(w http.ResponseWriter, r *http.Request) {
http.Redirect(w, r, "/two", http.StatusFound)
})
mux.HandleFunc("/two", func(w http.ResponseWriter, r *http.Request) {
http.Redirect(w, r, blocked, http.StatusFound)
})
server := httptest.NewServer(&mux)
t.Cleanup(server.Close)
_, err := testFetcher(t, nil).Fetch(context.Background(), server.URL+"/one")
if !errors.Is(err, ErrBlocked) {
t.Fatalf("err = %v, want ErrBlocked", err)
}
}
func TestFetchRedirectToUnsupportedScheme(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
http.Redirect(w, r, "file:///etc/passwd", http.StatusFound)
}))
t.Cleanup(server.Close)
_, err := testFetcher(t, nil).Fetch(context.Background(), server.URL)
if !errors.Is(err, ErrUnsupported) {
t.Fatalf("err = %v, want ErrUnsupported", err)
}
}
func TestFetchRedirectLimit(t *testing.T) {
var requests atomic.Int64
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
requests.Add(1)
http.Redirect(w, r, "/next", http.StatusFound)
}))
t.Cleanup(server.Close)
_, err := testFetcher(t, nil).Fetch(context.Background(), server.URL)
if !errors.Is(err, errTooManyRedirects) {
t.Fatalf("err = %v, want errTooManyRedirects", err)
}
if got := requests.Load(); got != maxRedirects+1 {
t.Errorf("запросов = %d, want %d (исходный + %d перехода)", got, maxRedirects+1, maxRedirects)
}
}
// TestFetchFinalURLIsBaseForRelativeLinks: относительная картинка считается от
// адреса после редиректов, а не от исходной ссылки.
func TestFetchFinalURLIsBaseForRelativeLinks(t *testing.T) {
var mux http.ServeMux
mux.HandleFunc("/start", func(w http.ResponseWriter, r *http.Request) {
http.Redirect(w, r, "/dir/final", http.StatusFound)
})
mux.HandleFunc("/dir/final", func(w http.ResponseWriter, _ *http.Request) {
w.Header().Set("Content-Type", "text/html")
_, _ = io.WriteString(w, `<html><head><meta property="og:image" content="pic.png"></head></html>`)
})
server := httptest.NewServer(&mux)
t.Cleanup(server.Close)
preview, err := testFetcher(t, nil).Fetch(context.Background(), server.URL+"/start")
if err != nil {
t.Fatalf("Fetch() вернул ошибку: %v", err)
}
if want := server.URL + "/dir/pic.png"; preview.ImageURL != want {
t.Errorf("ImageURL = %q, want %q", preview.ImageURL, want)
}
}
func TestFetchTimeout(t *testing.T) {
const timeout = 250 * time.Millisecond
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
time.Sleep(800 * time.Millisecond) // дольше таймаута загрузчика
w.Header().Set("Content-Type", "text/html")
_, _ = io.WriteString(w, `<html><head><title>поздно</title></head></html>`)
}))
t.Cleanup(server.Close)
fetcher := testFetcher(t, func(options *Options) { options.Timeout = timeout })
start := time.Now()
_, err := fetcher.Fetch(context.Background(), server.URL)
elapsed := time.Since(start)
if err == nil {
t.Fatal("Fetch() вернул nil, want ошибку таймаута")
}
if !errors.Is(err, context.DeadlineExceeded) && !os.IsTimeout(err) {
t.Errorf("err = %v, want таймаут", err)
}
if elapsed > 2*timeout {
t.Errorf("Fetch() занял %v, want не больше %v", elapsed, 2*timeout)
}
}
func TestFetchDropsBlockedImage(t *testing.T) {
cases := []struct {
name string
page string
}{
{"частный адрес", `<meta property="og:image" content="http://10.0.0.1/x.png">`},
{"метаданные облака", `<meta property="og:image" content="http://169.254.169.254/x.png">`},
{"javascript", `<meta property="og:image" content="javascript:alert(1)">`},
{"data", `<meta property="og:image" content="data:image/png;base64,AAAA">`},
{"localhost", `<meta property="og:image" content="http://localhost/x.png">`},
}
for _, testCase := range cases {
t.Run(testCase.name, func(t *testing.T) {
page := `<html><head><meta property="og:title" content="Заголовок">` + testCase.page + `</head></html>`
server := serveHTML(t, "text/html", page)
preview, err := testFetcher(t, nil).Fetch(context.Background(), server.URL)
if err != nil {
t.Fatalf("Fetch() вернул ошибку: %v", err)
}
if preview.ImageURL != "" {
t.Errorf("ImageURL = %q, want пусто", preview.ImageURL)
}
if preview.Title != "Заголовок" {
t.Errorf("Title = %q, want %q", preview.Title, "Заголовок")
}
})
}
}
func TestFetchSanitizesText(t *testing.T) {
const page = "<html><head>" +
"<meta property=\"og:title\" content=\"Заголовок\u0001 с \u0007управляющими\n\n\tсимволами и пробелами\">" +
"<meta property=\"og:description\" content=\" описание\nс переводом \">" +
"<meta property=\"og:site_name\" content=\" сайт \">" +
"</head></html>"
server := serveHTML(t, "text/html", page)
preview, err := testFetcher(t, nil).Fetch(context.Background(), server.URL)
if err != nil {
t.Fatalf("Fetch() вернул ошибку: %v", err)
}
if want := "Заголовок с управляющими символами и пробелами"; preview.Title != want {
t.Errorf("Title = %q, want %q", preview.Title, want)
}
if want := "описание с переводом"; preview.Description != want {
t.Errorf("Description = %q, want %q", preview.Description, want)
}
if want := "сайт"; preview.SiteName != want {
t.Errorf("SiteName = %q, want %q", preview.SiteName, want)
}
}
func TestFetchTruncatesLongText(t *testing.T) {
page := `<html><head>` +
`<meta property="og:title" content="` + strings.Repeat("я", 500) + `">` +
`<meta property="og:description" content="` + strings.Repeat("d", 900) + `">` +
`<meta property="og:site_name" content="` + strings.Repeat("s", 300) + `">` +
`</head></html>`
server := serveHTML(t, "text/html", page)
preview, err := testFetcher(t, nil).Fetch(context.Background(), server.URL)
if err != nil {
t.Fatalf("Fetch() вернул ошибку: %v", err)
}
if got := utf8.RuneCountInString(preview.Title); got != titleLimit {
t.Errorf("длина Title = %d рун, want %d", got, titleLimit)
}
if got := utf8.RuneCountInString(preview.Description); got != descriptionLimit {
t.Errorf("длина Description = %d рун, want %d", got, descriptionLimit)
}
if got := utf8.RuneCountInString(preview.SiteName); got != siteNameLimit {
t.Errorf("длина SiteName = %d рун, want %d", got, siteNameLimit)
}
}
func TestFetchPageWithoutMetadata(t *testing.T) {
server := serveHTML(t, "text/html", `<html><head><meta charset="utf-8"></head><body>текст</body></html>`)
preview, err := testFetcher(t, nil).Fetch(context.Background(), server.URL+"/page")
if err != nil {
t.Fatalf("Fetch() вернул ошибку: %v", err)
}
if !preview.Empty() {
t.Errorf("превью = %+v, want пустое", preview)
}
if want := server.URL + "/page"; preview.URL != want {
t.Errorf("URL = %q, want %q", preview.URL, want)
}
}
// TestFetchHidesQueryInErrorsAndLogs: полный URL с query не попадает ни в
// ошибку, ни в лог — только хост и код причины (приватность).
func TestFetchHidesQueryInErrorsAndLogs(t *testing.T) {
cases := []struct {
name string
rawURL string
}{
{"запрещённый литерал", "http://127.0.0.1:1/page?token=secret"},
{"ошибка соединения", "http://rebind.invalid/page?token=secret"},
}
for _, testCase := range cases {
t.Run(testCase.name, func(t *testing.T) {
var logged bytes.Buffer
logger := slog.New(slog.NewTextHandler(&logged, &slog.HandlerOptions{Level: slog.LevelDebug}))
fetcher := NewFetcher(Options{
Timeout: time.Second,
Logger: logger,
Resolver: fakeResolver(t, netip.MustParseAddr("10.1.2.3")),
})
_, err := fetcher.Fetch(context.Background(), testCase.rawURL)
if !errors.Is(err, ErrBlocked) {
t.Fatalf("err = %v, want ErrBlocked", err)
}
if strings.Contains(err.Error(), "secret") {
t.Errorf("в ошибке остался query: %v", err)
}
if strings.Contains(logged.String(), "secret") {
t.Errorf("в логе остался query: %s", logged.String())
}
if !strings.Contains(logged.String(), "host=") || !strings.Contains(logged.String(), "blocked") {
t.Errorf("лог не содержит хост и код причины: %s", logged.String())
}
})
}
}
// TestFetchBlocksDNSRebinding: имя резолвится в частный адрес, но соединение
// не открывается — адрес проверяется в момент подключения. Локальный
// DNS-сервер отдаёт 10.1.2.3, наружу тест не ходит.
func TestFetchBlocksDNSRebinding(t *testing.T) {
resolver := fakeResolver(t, netip.MustParseAddr("10.1.2.3"))
fetcher := NewFetcher(Options{Timeout: time.Second, Resolver: resolver})
start := time.Now()
_, err := fetcher.Fetch(context.Background(), "http://rebind.invalid/page")
if !errors.Is(err, ErrBlocked) {
t.Fatalf("err = %v, want ErrBlocked", err)
}
if elapsed := time.Since(start); elapsed > time.Second {
t.Errorf("проверка заняла %v — похоже, соединение всё-таки открывалось", elapsed)
}
}
// fakeResolver поднимает минимальный DNS-сервер на петле и отвечает указанным
// адресом на любой A-запрос. Нужен, чтобы проверить защиту от DNS rebinding
// без выхода во внешнюю сеть.
func fakeResolver(t *testing.T, answer netip.Addr) *net.Resolver {
t.Helper()
var listenConfig net.ListenConfig
conn, err := listenConfig.ListenPacket(context.Background(), "udp", "127.0.0.1:0")
if err != nil {
t.Fatalf("не удалось занять UDP-порт: %v", err)
}
t.Cleanup(func() { _ = conn.Close() })
go func() {
buffer := make([]byte, 1500)
for {
n, addr, err := conn.ReadFrom(buffer)
if err != nil {
return
}
if response := dnsAnswer(buffer[:n], answer); response != nil {
_, _ = conn.WriteTo(response, addr)
}
}
}()
return &net.Resolver{
PreferGo: true,
Dial: func(ctx context.Context, network, _ string) (net.Conn, error) {
protocol := "udp"
if strings.HasPrefix(network, "tcp") {
protocol = "tcp"
}
var dialer net.Dialer
return dialer.DialContext(ctx, protocol, conn.LocalAddr().String())
},
}
}
// dnsAnswer собирает ответ на запрос: A-запись с нужным адресом, для остальных
// типов — пустой список ответов (NODATA). Ошибки разбора означают, что пакет
// не наш — такой запрос просто игнорируется.
func dnsAnswer(query []byte, answer netip.Addr) []byte {
var parser dnsmessage.Parser
header, err := parser.Start(query)
if err != nil {
return nil
}
question, err := parser.Question()
if err != nil {
return nil
}
builder := dnsmessage.NewBuilder(nil, dnsmessage.Header{
ID: header.ID,
Response: true,
Authoritative: true,
RecursionAvailable: true,
})
builder.EnableCompression()
if err := builder.StartQuestions(); err != nil {
return nil
}
if err := builder.Question(question); err != nil {
return nil
}
if err := builder.StartAnswers(); err != nil {
return nil
}
if question.Type == dnsmessage.TypeA && answer.Is4() {
err := builder.AResource(
dnsmessage.ResourceHeader{Name: question.Name, Class: dnsmessage.ClassINET, TTL: 60},
dnsmessage.AResource{A: answer.As4()})
if err != nil {
return nil
}
}
message, err := builder.Finish()
if err != nil {
return nil
}
return message
}