gotextlog

gotextlog
git clone https://git.ryansepassi.com/git/gotextlog.git
Log | Files | Refs | README

client_test.go (17522B)


      1 package textlog
      2 
      3 import (
      4 	"context"
      5 	"crypto/ed25519"
      6 	"crypto/rand"
      7 	"encoding/json"
      8 	"errors"
      9 	"fmt"
     10 	"io"
     11 	"net/http"
     12 	"net/http/httptest"
     13 	"net/url"
     14 	"reflect"
     15 	"strings"
     16 	"testing"
     17 	"time"
     18 
     19 	"golang.org/x/crypto/ssh"
     20 )
     21 
     22 type expectedRequest struct {
     23 	method string
     24 	path   string
     25 	query  url.Values
     26 	body   string
     27 	reply  string
     28 }
     29 
     30 func TestClientAPIContract(t *testing.T) {
     31 	post := `{"id":9,"top_id":1,"body":"hi","created_at":"2026-01-01T00:00:00Z","parent_id":1,"reply_count":2,"tags":["go"],"mentions":["david"],"url":"https://text.test/@me/9","api_url":"https://text.test/api/v1/posts/9","author":{"handle":"me","url":"https://text.test/@me","api_url":"https://text.test/api/v1/users/me"}}`
     32 	collection := `{"data":[],"pagination":{"next_cursor":null}}`
     33 	expectations := []expectedRequest{
     34 		{http.MethodGet, "/api/v1/feeds/latest", values("limit", "2", "cursor", "next one"), "", collection},
     35 		{http.MethodGet, "/api/v1/feeds/hot", values("limit", "3"), "", collection},
     36 		{http.MethodGet, "/api/v1/activities/for-you", values("limit", "4", "cursor", "c"), "", `{"data":[],"pagination":{"next_cursor":"n"},"has_unread":true}`},
     37 		{http.MethodGet, "/api/v1/activities/to-me", values("limit", "5"), "", `{"data":[],"pagination":{"next_cursor":null},"has_unread":false}`},
     38 		{http.MethodGet, "/api/v1/search", values("q", "quiet thoughts", "limit", "6", "cursor", "s"), "", collection},
     39 		{http.MethodGet, "/api/v1/posts/9", nil, "", `{"data":` + post + `}`},
     40 		{http.MethodGet, "/api/v1/posts/9/replies", values("depth", "5", "limit", "100", "cursor", "r"), "", collection},
     41 		{http.MethodGet, "/api/v1/users/a%2Fb%20c", nil, "", `{"data":{"handle":"a/b c","bio":"bio","created_at":"2026-01-01T00:00:00Z","post_count":1,"replies_count":2,"follower_count":3,"following_user_count":4,"following_tag_count":5,"following_count":9,"blocked_user_count":6,"blocked_tag_count":7,"url":"u","api_url":"a"}}`},
     42 		{http.MethodGet, "/api/v1/users/me/notes", values("limit", "7", "cursor", "n"), "", collection},
     43 		{http.MethodGet, "/api/v1/users/me/posts", values("limit", "8"), "", collection},
     44 		{http.MethodGet, "/api/v1/users/me/replies", values("limit", "9", "cursor", "r"), "", collection},
     45 		{http.MethodGet, "/api/v1/users/me/following/users", values("limit", "10"), "", collection},
     46 		{http.MethodGet, "/api/v1/users/me/following/tags", values("limit", "11", "cursor", "t"), "", collection},
     47 		{http.MethodGet, "/api/v1/users/me/followers", values("limit", "12"), "", collection},
     48 		{http.MethodGet, "/api/v1/users/me/blocks", values("limit", "13", "cursor", "b"), "", collection},
     49 		{http.MethodGet, "/api/v1/tags/go%2Flang", nil, "", `{"data":{"tag":"go/lang","post_count":1,"follower_count":2,"url":"u","api_url":"a"}}`},
     50 		{http.MethodGet, "/api/v1/tags/go/posts", values("limit", "14", "cursor", "p"), "", collection},
     51 		{http.MethodGet, "/api/v1/tags/go/followers", values("limit", "15"), "", collection},
     52 		{http.MethodPost, "/api/v1/auth/request", nil, `{"email":"me@example.test"}`, `{"data":{"sent":true}}`},
     53 		{http.MethodPost, "/api/v1/auth/verify", nil, `{"code":"123456","email":"me@example.test"}`, `{"data":{"token":"new-token","expires_at":"2026-02-01T00:00:00Z","user":{"handle":"me","email":"me@example.test","bio":"bio","email_verified":true,"can_post":true}}}`},
     54 		{http.MethodDelete, "/api/v1/auth/session", nil, "", `{"data":{"revoked":true}}`},
     55 		{http.MethodGet, "/api/v1/me", nil, "", `{"data":{"handle":"me","email":"me@example.test","bio":"bio","email_verified":true,"can_post":true}}`},
     56 		{http.MethodPatch, "/api/v1/me", nil, `{"bio":"new bio"}`, `{"data":{"handle":"me","email":"me@example.test","bio":"new bio","email_verified":true,"can_post":true}}`},
     57 		{http.MethodPost, "/api/v1/posts", nil, `{"body":"hi","parent_id":1}`, `{"data":` + post + `}`},
     58 		{http.MethodPatch, "/api/v1/posts/9", nil, `{"body":"edited"}`, `{"data":` + post + `}`},
     59 		{http.MethodDelete, "/api/v1/posts/9", nil, "", `{"data":{"deleted":true}}`},
     60 		{http.MethodPost, "/api/v1/users/a%2Fb/follow", nil, "", `{"data":{"following":true}}`},
     61 		{http.MethodDelete, "/api/v1/users/a%2Fb/follow", nil, "", `{"data":{"following":false}}`},
     62 		{http.MethodPost, "/api/v1/users/a%2Fb/block", nil, "", `{"data":{"blocked":true}}`},
     63 		{http.MethodDelete, "/api/v1/users/a%2Fb/block", nil, "", `{"data":{"blocked":false}}`},
     64 		{http.MethodPost, "/api/v1/posts/9/report", nil, `{"reason":"spam"}`, `{"data":{"reported":true}}`},
     65 	}
     66 
     67 	requestIndex := 0
     68 	server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
     69 		if requestIndex >= len(expectations) {
     70 			t.Errorf("unexpected request %s %s", r.Method, r.URL.RequestURI())
     71 			http.Error(w, "unexpected request", http.StatusInternalServerError)
     72 			return
     73 		}
     74 		want := expectations[requestIndex]
     75 		requestIndex++
     76 		if got := r.Method; got != want.method {
     77 			t.Errorf("request %d method = %q, want %q", requestIndex, got, want.method)
     78 		}
     79 		if got := r.URL.EscapedPath(); got != want.path {
     80 			t.Errorf("request %d path = %q, want %q", requestIndex, got, want.path)
     81 		}
     82 		if got := r.URL.Query(); len(got) != len(want.query) || (len(got) > 0 && !reflect.DeepEqual(got, want.query)) {
     83 			t.Errorf("request %d query = %#v, want %#v", requestIndex, got, want.query)
     84 		}
     85 		if got := r.Header.Get("Accept"); got != "application/json" {
     86 			t.Errorf("request %d Accept = %q", requestIndex, got)
     87 		}
     88 		if got := r.Header.Get("Authorization"); got != "Bearer secret" {
     89 			t.Errorf("request %d Authorization = %q", requestIndex, got)
     90 		}
     91 		body, err := io.ReadAll(r.Body)
     92 		if err != nil {
     93 			t.Errorf("request %d body: %v", requestIndex, err)
     94 		}
     95 		if got := strings.TrimSpace(string(body)); got != want.body {
     96 			t.Errorf("request %d body = %q, want %q", requestIndex, got, want.body)
     97 		}
     98 		wantContentType := ""
     99 		if want.body != "" {
    100 			wantContentType = "application/json"
    101 		}
    102 		if got := r.Header.Get("Content-Type"); got != wantContentType {
    103 			t.Errorf("request %d Content-Type = %q, want %q", requestIndex, got, wantContentType)
    104 		}
    105 		w.Header().Set("Content-Type", "application/json")
    106 		fmt.Fprint(w, want.reply)
    107 	}))
    108 	defer server.Close()
    109 
    110 	ctx := context.Background()
    111 	api := NewClient(server.URL+"/ignored/base/path", "secret", server.Client())
    112 	_, err := api.Latest(ctx, 2, "next one")
    113 	mustSucceed(t, err)
    114 	_, err = api.Hot(ctx, 3, "")
    115 	mustSucceed(t, err)
    116 	forYou, err := api.ForYou(ctx, 4, "c")
    117 	mustSucceed(t, err)
    118 	if !forYou.HasUnread || forYou.Pagination.NextCursor == nil || *forYou.Pagination.NextCursor != "n" {
    119 		t.Fatalf("for-you response = %#v", forYou)
    120 	}
    121 	_, err = api.ToMe(ctx, 5, "")
    122 	mustSucceed(t, err)
    123 	_, err = api.Search(ctx, "quiet thoughts", 6, "s")
    124 	mustSucceed(t, err)
    125 	gotPost, err := api.Post(ctx, 9)
    126 	mustSucceed(t, err)
    127 	if gotPost.Data.ID != 9 || !gotPost.Data.CreatedAt.Equal(time.Date(2026, 1, 1, 0, 0, 0, 0, time.UTC)) {
    128 		t.Fatalf("post response = %#v", gotPost)
    129 	}
    130 	_, err = api.Replies(ctx, 9, 5, 100, "r")
    131 	mustSucceed(t, err)
    132 	_, err = api.Profile(ctx, "a/b c")
    133 	mustSucceed(t, err)
    134 	_, err = api.UserNotes(ctx, "me", 7, "n")
    135 	mustSucceed(t, err)
    136 	_, err = api.UserPosts(ctx, "me", 8, "")
    137 	mustSucceed(t, err)
    138 	_, err = api.UserReplies(ctx, "me", 9, "r")
    139 	mustSucceed(t, err)
    140 	_, err = api.UserFollowing(ctx, "me", 10, "")
    141 	mustSucceed(t, err)
    142 	_, err = api.UserFollowingTags(ctx, "me", 11, "t")
    143 	mustSucceed(t, err)
    144 	_, err = api.UserFollowers(ctx, "me", 12, "")
    145 	mustSucceed(t, err)
    146 	_, err = api.UserBlocks(ctx, "me", 13, "b")
    147 	mustSucceed(t, err)
    148 	_, err = api.Tag(ctx, "go/lang")
    149 	mustSucceed(t, err)
    150 	_, err = api.TagPosts(ctx, "go", 14, "p")
    151 	mustSucceed(t, err)
    152 	_, err = api.TagFollowers(ctx, "go", 15, "")
    153 	mustSucceed(t, err)
    154 	sent, err := api.RequestCode(ctx, "me@example.test")
    155 	mustSucceed(t, err)
    156 	if !sent.Data.Sent {
    157 		t.Fatal("request-code response did not decode")
    158 	}
    159 	session, err := api.VerifyCode(ctx, "me@example.test", "123456")
    160 	mustSucceed(t, err)
    161 	if session.Data.Token != "new-token" || session.Data.User.Handle != "me" {
    162 		t.Fatalf("verify response = %#v", session)
    163 	}
    164 	revoked, err := api.Revoke(ctx)
    165 	mustSucceed(t, err)
    166 	if !revoked.Data.Revoked {
    167 		t.Fatal("revoke response did not decode")
    168 	}
    169 	_, err = api.Me(ctx)
    170 	mustSucceed(t, err)
    171 	_, err = api.UpdateBio(ctx, "new bio")
    172 	mustSucceed(t, err)
    173 	parentID := 1
    174 	_, err = api.CreatePost(ctx, "hi", &parentID)
    175 	mustSucceed(t, err)
    176 	_, err = api.EditPost(ctx, 9, "edited")
    177 	mustSucceed(t, err)
    178 	deleted, err := api.DeletePost(ctx, 9)
    179 	mustSucceed(t, err)
    180 	if !deleted.Data.Deleted {
    181 		t.Fatal("delete response did not decode")
    182 	}
    183 	following, err := api.Follow(ctx, "a/b", true)
    184 	mustSucceed(t, err)
    185 	if !following.Data.Following {
    186 		t.Fatal("follow response did not decode")
    187 	}
    188 	_, err = api.Follow(ctx, "a/b", false)
    189 	mustSucceed(t, err)
    190 	blocked, err := api.Block(ctx, "a/b", true)
    191 	mustSucceed(t, err)
    192 	if !blocked.Data.Blocked {
    193 		t.Fatal("block response did not decode")
    194 	}
    195 	_, err = api.Block(ctx, "a/b", false)
    196 	mustSucceed(t, err)
    197 	reported, err := api.Report(ctx, 9, ReportSpam)
    198 	mustSucceed(t, err)
    199 	if !reported.Data.Reported {
    200 		t.Fatal("report response did not decode")
    201 	}
    202 
    203 	if requestIndex != len(expectations) {
    204 		t.Fatalf("received %d requests, want %d", requestIndex, len(expectations))
    205 	}
    206 }
    207 
    208 func TestClientErrorsAndContext(t *testing.T) {
    209 	t.Run("structured HTTP error", func(t *testing.T) {
    210 		server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
    211 			w.Header().Set("Retry-After", "12")
    212 			w.WriteHeader(http.StatusTooManyRequests)
    213 			fmt.Fprint(w, `{"error":{"code":"rate_limited","message":"Slow down"}}`)
    214 		}))
    215 		defer server.Close()
    216 
    217 		_, err := NewClient(server.URL, "", server.Client()).Latest(context.Background(), 20, "")
    218 		var apiErr *APIError
    219 		if !errors.As(err, &apiErr) {
    220 			t.Fatalf("error = %T %v, want *APIError", err, err)
    221 		}
    222 		if apiErr.Code != "rate_limited" || apiErr.Message != "Slow down" || apiErr.Status != http.StatusTooManyRequests || apiErr.RetryAfter != 12 {
    223 			t.Fatalf("API error = %#v", apiErr)
    224 		}
    225 		if got := ErrorMessage(err); got != "Slow down (retry in 12s)" {
    226 			t.Fatalf("ErrorMessage = %q", got)
    227 		}
    228 	})
    229 
    230 	t.Run("invalid error response", func(t *testing.T) {
    231 		server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
    232 			w.WriteHeader(http.StatusBadGateway)
    233 			fmt.Fprint(w, "not json")
    234 		}))
    235 		defer server.Close()
    236 
    237 		_, err := NewClient(server.URL, "", server.Client()).Hot(context.Background(), 20, "")
    238 		var apiErr *APIError
    239 		if !errors.As(err, &apiErr) || apiErr.Code != "http_error" || apiErr.Message != "HTTP 502" {
    240 			t.Fatalf("error = %#v", err)
    241 		}
    242 	})
    243 
    244 	t.Run("network and cancellation errors", func(t *testing.T) {
    245 		api := NewClient("http://example.test", "", &http.Client{Transport: roundTripFunc(func(*http.Request) (*http.Response, error) {
    246 			return nil, errors.New("dial failed")
    247 		})})
    248 		_, err := api.Me(context.Background())
    249 		var apiErr *APIError
    250 		if !errors.As(err, &apiErr) || apiErr.Code != "network_error" || apiErr.Status != 0 {
    251 			t.Fatalf("network error = %#v", err)
    252 		}
    253 
    254 		ctx, cancel := context.WithCancel(context.Background())
    255 		cancel()
    256 		_, err = NewClient("http://example.test", "", http.DefaultClient).Me(ctx)
    257 		if !errors.Is(err, context.Canceled) {
    258 			t.Fatalf("canceled request error = %T %v", err, err)
    259 		}
    260 	})
    261 }
    262 
    263 func TestFirehoseParsesChunkedSSEAndSkipsMalformedEvents(t *testing.T) {
    264 	server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
    265 		if got := r.Header.Get("Accept"); got != "text/event-stream" {
    266 			t.Errorf("Accept = %q", got)
    267 		}
    268 		if got := r.Header.Get("Authorization"); got != "Bearer secret" {
    269 			t.Errorf("Authorization = %q", got)
    270 		}
    271 		w.Header().Set("Content-Type", "text/event-stream")
    272 		flusher := w.(http.Flusher)
    273 		fmt.Fprint(w, ": welcome\r\nevent: ready\r\ndata: {}\r\n\r\nevent: post\r\ndata: {\"id\":1,")
    274 		flusher.Flush()
    275 		fmt.Fprint(w, "\"body\":\"first\",\r\ndata: \"created_at\":\"2026-01-01T00:00:00Z\",\"author\":{\"handle\":\"david\"}}\r\n\r\n")
    276 		fmt.Fprint(w, "event: post\ndata: {bad}\n\nevent: post\ndata: {\"data\":{\"id\":2,\"body\":\"second\",\"created_at\":\"2026-01-02T00:00:00Z\",\"author\":{\"handle\":\"amy\"}}}\n\n")
    277 		fmt.Fprint(w, "data: {\"id\":3,\"body\":\"wrong event\",\"created_at\":\"2026-01-03T00:00:00Z\",\"author\":{\"handle\":\"x\"}}\n\n")
    278 		fmt.Fprint(w, "event: post\ndata: {\"id\":4,\"body\":\"final\",\"created_at\":\"2026-01-04T00:00:00Z\",\"author\":{\"handle\":\"zoe\"}}")
    279 	}))
    280 	defer server.Close()
    281 
    282 	posts, errs := NewClient(server.URL, "secret", server.Client()).Firehose(context.Background())
    283 	var got []Post
    284 	for post := range posts {
    285 		got = append(got, post)
    286 	}
    287 	if err := <-errs; err != nil {
    288 		t.Fatal(err)
    289 	}
    290 	if len(got) != 3 {
    291 		t.Fatalf("posts = %#v", got)
    292 	}
    293 	if gotIDs := []int{got[0].ID, got[1].ID, got[2].ID}; !reflect.DeepEqual(gotIDs, []int{1, 2, 4}) {
    294 		t.Fatalf("post IDs = %v", gotIDs)
    295 	}
    296 }
    297 
    298 func TestFirehoseReportsHTTPAndCancellationErrors(t *testing.T) {
    299 	t.Run("HTTP error", func(t *testing.T) {
    300 		server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
    301 			w.WriteHeader(http.StatusServiceUnavailable)
    302 		}))
    303 		defer server.Close()
    304 
    305 		posts, errs := NewClient(server.URL, "", server.Client()).Firehose(context.Background())
    306 		if _, ok := <-posts; ok {
    307 			t.Fatal("unexpected post")
    308 		}
    309 		var apiErr *APIError
    310 		if err := <-errs; !errors.As(err, &apiErr) || apiErr.Code != "stream_error" || apiErr.Status != http.StatusServiceUnavailable {
    311 			t.Fatalf("stream error = %#v", err)
    312 		}
    313 	})
    314 
    315 	t.Run("cancellation", func(t *testing.T) {
    316 		started := make(chan struct{})
    317 		server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
    318 			w.Header().Set("Content-Type", "text/event-stream")
    319 			w.WriteHeader(http.StatusOK)
    320 			w.(http.Flusher).Flush()
    321 			close(started)
    322 			<-r.Context().Done()
    323 		}))
    324 		defer server.Close()
    325 
    326 		ctx, cancel := context.WithCancel(context.Background())
    327 		posts, errs := NewClient(server.URL, "", server.Client()).Firehose(ctx)
    328 		<-started
    329 		cancel()
    330 		for range posts {
    331 		}
    332 		if err := <-errs; !errors.Is(err, context.Canceled) {
    333 			t.Fatalf("cancellation error = %T %v", err, err)
    334 		}
    335 	})
    336 }
    337 
    338 func TestClientKeyAuthenticationExchange(t *testing.T) {
    339 	publicRaw, privateKey, err := ed25519.GenerateKey(rand.Reader)
    340 	mustSucceed(t, err)
    341 	publicKey, err := ssh.NewPublicKey(publicRaw)
    342 	mustSucceed(t, err)
    343 	signer, err := ssh.NewSignerFromKey(privateKey)
    344 	mustSucceed(t, err)
    345 
    346 	message := []byte("gotextlog-auth-v1\nhttps://text.test\nchallenge")
    347 	expiresAt := time.Now().UTC().Add(time.Minute).Truncate(time.Second)
    348 	server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) {
    349 		writer.Header().Set("Content-Type", "application/json")
    350 		switch request.URL.Path {
    351 		case "/api/v1/auth/key/challenge":
    352 			var body struct {
    353 				PublicKey string `json:"public_key"`
    354 				Handle    string `json:"handle"`
    355 			}
    356 			if request.Method != http.MethodPost || json.NewDecoder(request.Body).Decode(&body) != nil {
    357 				t.Errorf("invalid challenge request")
    358 			}
    359 			parsed, _, _, _, parseErr := ssh.ParseAuthorizedKey([]byte(body.PublicKey))
    360 			if parseErr != nil || string(parsed.Marshal()) != string(publicKey.Marshal()) || body.Handle != "alice" {
    361 				t.Errorf("challenge request = %#v, parse error = %v", body, parseErr)
    362 			}
    363 			_ = json.NewEncoder(writer).Encode(Envelope[KeyChallenge]{Data: KeyChallenge{
    364 				ChallengeID: "challenge-1", Message: message, ExpiresAt: expiresAt,
    365 			}})
    366 		case "/api/v1/auth/key/verify":
    367 			var body struct {
    368 				ChallengeID string `json:"challenge_id"`
    369 				Signature   []byte `json:"signature"`
    370 			}
    371 			if request.Method != http.MethodPost || json.NewDecoder(request.Body).Decode(&body) != nil {
    372 				t.Errorf("invalid verify request")
    373 			}
    374 			var signature ssh.Signature
    375 			if body.ChallengeID != "challenge-1" || ssh.Unmarshal(body.Signature, &signature) != nil || publicKey.Verify(message, &signature) != nil {
    376 				t.Errorf("signature did not verify")
    377 			}
    378 			_ = json.NewEncoder(writer).Encode(Envelope[Session]{Data: Session{
    379 				Token: "session-token", ExpiresAt: expiresAt, User: CurrentUser{Handle: "alice", CanPost: true},
    380 			}})
    381 		default:
    382 			http.NotFound(writer, request)
    383 		}
    384 	}))
    385 	defer server.Close()
    386 
    387 	client := NewClient(server.URL, "", server.Client())
    388 	challenge, err := client.RequestKeyChallenge(context.Background(), strings.TrimSpace(string(ssh.MarshalAuthorizedKey(publicKey))), "alice")
    389 	mustSucceed(t, err)
    390 	if challenge.Data.ChallengeID != "challenge-1" || !reflect.DeepEqual(challenge.Data.Message, message) {
    391 		t.Fatalf("challenge = %#v", challenge.Data)
    392 	}
    393 	signature, err := signer.Sign(rand.Reader, challenge.Data.Message)
    394 	mustSucceed(t, err)
    395 	session, err := client.VerifyKey(context.Background(), challenge.Data.ChallengeID, ssh.Marshal(signature))
    396 	mustSucceed(t, err)
    397 	if session.Data.Token != "session-token" || session.Data.User.Handle != "alice" {
    398 		t.Fatalf("session = %#v", session.Data)
    399 	}
    400 }
    401 
    402 func values(entries ...string) url.Values {
    403 	result := make(url.Values, len(entries)/2)
    404 	for index := 0; index < len(entries); index += 2 {
    405 		result.Set(entries[index], entries[index+1])
    406 	}
    407 	return result
    408 }
    409 
    410 func mustSucceed(t *testing.T, err error) {
    411 	t.Helper()
    412 	if err != nil {
    413 		t.Fatal(err)
    414 	}
    415 }
    416 
    417 type roundTripFunc func(*http.Request) (*http.Response, error)
    418 
    419 func (fn roundTripFunc) RoundTrip(request *http.Request) (*http.Response, error) {
    420 	return fn(request)
    421 }
    422 
    423 var _ http.RoundTripper = roundTripFunc(nil)