package api import ( "encoding/json" "net/http" "net/http/httptest" "strings" "testing" "time" "oneblog/internal/auth" "oneblog/internal/config" "oneblog/internal/db" "oneblog/internal/model" "oneblog/internal/store" ) func newTestAPI(t *testing.T) (*API, http.Handler) { t.Helper() d, err := db.Open("sqlite", ":memory:") if err != nil { t.Fatalf("open sqlite: %v", err) } t.Cleanup(func() { d.Close() }) st, err := store.New(d) if err != nil { t.Fatalf("store: %v", err) } a := &API{ Store: st, Cfg: &config.Config{SiteURL: "http://localhost:8080"}, ReaderSessions: auth.NewReaderSessions("test-secret", time.Hour), TG: auth.Telegram{Bot: "testbot", Token: "123:abc"}, } settings, err := st.GetSettings() if err != nil { t.Fatalf("settings: %v", err) } settings.CommentsEnabled = true if err := st.UpdateSettings(settings); err != nil { t.Fatalf("enable comments: %v", err) } return a, a.Routes() } // Telegram 伪造签名反复重试应触发 IP 限速 func TestTelegramAuthRateLimited(t *testing.T) { _, h := newTestAPI(t) body := `{"id":1,"first_name":"x","hash":"deadbeef"}` var rec *httptest.ResponseRecorder for i := 0; i < maxAuthFails; i++ { req := httptest.NewRequest(http.MethodPost, "/api/auth/telegram", strings.NewReader(body)) rec = httptest.NewRecorder() h.ServeHTTP(rec, req) if rec.Code != http.StatusForbidden { t.Fatalf("attempt %d: got %d, want 403", i+1, rec.Code) } } req := httptest.NewRequest(http.MethodPost, "/api/auth/telegram", strings.NewReader(body)) rec = httptest.NewRecorder() h.ServeHTTP(rec, req) if rec.Code != http.StatusTooManyRequests { t.Fatalf("after %d failures: got %d, want 429", maxAuthFails, rec.Code) } } // 评论写入按读者限速:第 maxComments+1 条被拒 func TestCommentRateLimited(t *testing.T) { a, h := newTestAPI(t) p, err := a.Store.Create(model.PostInput{Kind: model.KindLong, Title: "t", Slug: "t", ContentMd: "x", Status: model.StatusPublished}) if err != nil { t.Fatal(err) } rd, err := a.Store.UpsertReader(model.Reader{Provider: "github", Handle: "u1", Name: "u1"}) if err != nil { t.Fatal(err) } tok, _ := a.ReaderSessions.Issue(rd.ID) post := func(i int) *httptest.ResponseRecorder { body, _ := json.Marshal(map[string]any{"post_id": p.ID, "body_md": "好"}) req := httptest.NewRequest(http.MethodPost, "/api/comments", strings.NewReader(string(body))) req.AddCookie(&http.Cookie{Name: auth.ReaderCookie, Value: tok}) rec := httptest.NewRecorder() h.ServeHTTP(rec, req) return rec } for i := 0; i < maxComments; i++ { if rec := post(i); rec.Code != http.StatusCreated { t.Fatalf("comment %d: got %d %s", i+1, rec.Code, rec.Body.String()) } } if rec := post(maxComments); rec.Code != http.StatusTooManyRequests { t.Fatalf("over limit: got %d, want 429", rec.Code) } }