summaryrefslogtreecommitdiffstats
path: root/internal/nntp/client.go
diff options
context:
space:
mode:
Diffstat (limited to 'internal/nntp/client.go')
-rw-r--r--internal/nntp/client.go983
1 files changed, 983 insertions, 0 deletions
diff --git a/internal/nntp/client.go b/internal/nntp/client.go
new file mode 100644
index 0000000..5958628
--- /dev/null
+++ b/internal/nntp/client.go
@@ -0,0 +1,983 @@
+package nntp
+
+import (
+ "bufio"
+ "compress/flate"
+ "context"
+ "crypto/tls"
+ "encoding/base64"
+ "errors"
+ "fmt"
+ "io"
+ "net"
+ "strconv"
+ "strings"
+ "sync"
+ "time"
+ "unicode/utf8"
+
+ "golang.org/x/net/proxy"
+)
+
+const (
+ defaultTimeout = 30 * time.Second
+ maxResponseLine = 1 << 20
+ maxMultilineBytes = 64 << 20
+ maxMultilineLines = 1_000_000
+ maxPostBytes = 10 << 20
+ maxPostLineBytes = 998
+)
+
+var tls12AESGCMSuites = []uint16{
+ tls.TLS_ECDHE_RSA_WITH_AES_128_GCM_SHA256,
+ tls.TLS_ECDHE_RSA_WITH_AES_256_GCM_SHA384,
+ tls.TLS_ECDHE_ECDSA_WITH_AES_128_GCM_SHA256,
+ tls.TLS_ECDHE_ECDSA_WITH_AES_256_GCM_SHA384,
+}
+
+type DialConfig struct {
+ Host string
+ Port string
+ UseTLS bool
+ StartTLS bool
+ InsecureSkipVerify bool
+ Username string
+ Password string
+ SASLMechanism string
+ UseCompression bool
+ ProxyType string
+ ProxyAddress string
+ Timeout time.Duration
+ tlsConfig *tls.Config
+}
+
+type GroupInfo struct {
+ Name string
+ Low int64
+ High int64
+ EstimatedPost int64
+ Posting string
+}
+
+type GroupStatus struct {
+ Name string
+ Count int64
+ Low int64
+ High int64
+}
+
+type ArticleHeader struct {
+ Number int64
+ Subject string
+ From string
+ Date string
+ MessageID string
+ References string
+ Bytes int64
+ Lines int64
+}
+
+type HeaderValue struct {
+ Number int64
+ Value string
+}
+
+type GroupDescription struct {
+ Name string
+ Description string
+}
+
+type ServerDate struct {
+ Date string
+ Time string
+}
+
+type ResponseError struct {
+ Code int
+ Message string
+}
+
+func (e *ResponseError) Error() string {
+ return fmt.Sprintf("NNTP response %d: %s", e.Code, e.Message)
+}
+
+type Client struct {
+ mu sync.Mutex
+ conn net.Conn
+ reader *bufio.Reader
+ writer *bufio.Writer
+ timeout time.Duration
+ closed bool
+}
+
+func Dial(ctx context.Context, cfg DialConfig) (*Client, error) {
+ if err := validateDialConfig(cfg); err != nil {
+ return nil, err
+ }
+ if cfg.UseCompression && cfg.UseTLS {
+ return nil, errors.New("COMPRESS DEFLATE cannot be used with TLS")
+ }
+ timeout := cfg.Timeout
+ if timeout <= 0 {
+ timeout = defaultTimeout
+ }
+ target := net.JoinHostPort(cfg.Host, cfg.Port)
+ netDialer := &net.Dialer{Timeout: timeout, KeepAlive: 30 * time.Second}
+ var raw net.Conn
+ var err error
+ if strings.EqualFold(cfg.ProxyType, "SOCKS5") {
+ proxyDialer, proxyErr := proxy.SOCKS5("tcp", cfg.ProxyAddress, nil, netDialer)
+ if proxyErr != nil {
+ return nil, fmt.Errorf("configure SOCKS5 proxy: %w", proxyErr)
+ }
+ raw, err = proxyDialer.Dial("tcp", target)
+ } else {
+ raw, err = netDialer.DialContext(ctx, "tcp", target)
+ }
+ if err != nil {
+ return nil, fmt.Errorf("connect to NNTP server: %w", err)
+ }
+ conn := raw
+ if cfg.UseTLS && !cfg.StartTLS {
+ tlsConfig := makeTLSConfig(cfg)
+ tlsConn := tls.Client(raw, tlsConfig)
+ if err := tlsConn.HandshakeContext(ctx); err != nil {
+ raw.Close()
+ return nil, fmt.Errorf("complete TLS handshake: %w", err)
+ }
+ conn = tlsConn
+ }
+ client := &Client{
+ conn: conn,
+ reader: bufio.NewReaderSize(conn, 64*1024),
+ writer: bufio.NewWriterSize(conn, 64*1024),
+ timeout: timeout,
+ }
+ if err := client.setDeadline(); err != nil {
+ client.Close()
+ return nil, err
+ }
+ code, message, err := client.readResponse()
+ if err != nil {
+ client.Close()
+ return nil, fmt.Errorf("read NNTP greeting: %w", err)
+ }
+ if code != 200 && code != 201 {
+ client.Close()
+ return nil, &ResponseError{Code: code, Message: message}
+ }
+ if cfg.StartTLS {
+ code, message, err = client.command("STARTTLS")
+ if err != nil {
+ client.Close()
+ return nil, err
+ }
+ if code != 382 {
+ client.Close()
+ return nil, &ResponseError{Code: code, Message: message}
+ }
+ tlsConn := tls.Client(client.conn, makeTLSConfig(cfg))
+ if err := tlsConn.HandshakeContext(ctx); err != nil {
+ _ = client.conn.Close()
+ return nil, fmt.Errorf("complete STARTTLS handshake: %w", err)
+ }
+ client.conn = tlsConn
+ client.reader = bufio.NewReaderSize(tlsConn, 64*1024)
+ client.writer = bufio.NewWriterSize(tlsConn, 64*1024)
+ if err := client.setDeadline(); err != nil {
+ client.Close()
+ return nil, err
+ }
+ }
+ if cfg.Username != "" {
+ if err := client.authenticateWithConfig(cfg); err != nil {
+ client.Close()
+ return nil, err
+ }
+ }
+ if cfg.UseCompression {
+ if err := client.enableCompression(); err != nil {
+ client.Close()
+ return nil, err
+ }
+ }
+ if code, message, err = client.command("MODE READER"); err != nil {
+ client.Close()
+ return nil, err
+ }
+ if code != 200 && code != 201 && code != 500 && code != 501 {
+ client.Close()
+ return nil, &ResponseError{Code: code, Message: message}
+ }
+ _ = client.conn.SetDeadline(time.Time{})
+ return client, nil
+}
+
+func (c *Client) Capabilities() ([]string, error) {
+ c.mu.Lock()
+ defer c.mu.Unlock()
+ lines, err := c.capabilitiesUnlocked()
+ if responseCode(err) == 500 || responseCode(err) == 501 {
+ return nil, nil
+ }
+ return lines, err
+}
+
+func (c *Client) capabilitiesUnlocked() ([]string, error) {
+ return c.multilineCommand("CAPABILITIES", 101)
+}
+
+func (c *Client) Help() ([]string, error) {
+ c.mu.Lock()
+ defer c.mu.Unlock()
+ return c.multilineCommand("HELP", 100)
+}
+
+func (c *Client) Date() (ServerDate, error) {
+ c.mu.Lock()
+ defer c.mu.Unlock()
+ code, message, err := c.commandUnlocked("DATE")
+ if err != nil {
+ return ServerDate{}, err
+ }
+ if code != 111 {
+ return ServerDate{}, &ResponseError{Code: code, Message: message}
+ }
+ fields := strings.Fields(message)
+ if len(fields) < 2 {
+ return ServerDate{}, errors.New("malformed DATE response")
+ }
+ return ServerDate{Date: fields[0], Time: fields[1]}, nil
+}
+
+func (c *Client) ListOverviewFormat() ([]string, error) {
+ c.mu.Lock()
+ defer c.mu.Unlock()
+ return c.multilineCommand("LIST OVERVIEW.FMT", 215)
+}
+
+func (c *Client) ListNewsGroups(pattern string) ([]GroupDescription, error) {
+ if strings.ContainsAny(pattern, "\r\n") {
+ return nil, errors.New("invalid LIST NEWSGROUPS pattern")
+ }
+ command := "LIST NEWSGROUPS"
+ if strings.TrimSpace(pattern) != "" {
+ command += " " + strings.TrimSpace(pattern)
+ }
+ c.mu.Lock()
+ defer c.mu.Unlock()
+ lines, err := c.multilineCommand(command, 215)
+ if err != nil {
+ return nil, err
+ }
+ groups := make([]GroupDescription, 0, len(lines))
+ for _, line := range lines {
+ fields := strings.SplitN(line, " ", 2)
+ if len(fields) == 0 || !validAtom(fields[0]) {
+ continue
+ }
+ description := ""
+ if len(fields) == 2 {
+ description = strings.TrimSpace(fields[1])
+ }
+ groups = append(groups, GroupDescription{Name: fields[0], Description: description})
+ }
+ return groups, nil
+}
+
+func (c *Client) NewGroups(date, clock, timezone string) ([]string, error) {
+ if !validCommandValue(date) || !validCommandValue(clock) || !validCommandValue(timezone) {
+ return nil, errors.New("invalid NEWGROUPS arguments")
+ }
+ c.mu.Lock()
+ defer c.mu.Unlock()
+ return c.multilineCommand("NEWGROUPS "+date+" "+clock+" "+timezone, 231)
+}
+
+func (c *Client) NewNews(groups, date, clock, timezone string) ([]string, error) {
+ if !validCommandValue(groups) || !validCommandValue(date) || !validCommandValue(clock) || !validCommandValue(timezone) {
+ return nil, errors.New("invalid NEWNEWS arguments")
+ }
+ c.mu.Lock()
+ defer c.mu.Unlock()
+ return c.multilineCommand("NEWNEWS "+groups+" "+date+" "+clock+" "+timezone, 230)
+}
+
+func (c *Client) ListGroup(group string) ([]int64, error) {
+ if !validAtom(group) {
+ return nil, errors.New("invalid newsgroup name")
+ }
+ c.mu.Lock()
+ defer c.mu.Unlock()
+ lines, err := c.multilineCommand("LISTGROUP "+group, 211)
+ if err != nil {
+ return nil, err
+ }
+ articles := make([]int64, 0, len(lines))
+ for _, line := range lines {
+ for _, token := range strings.Fields(line) {
+ parts := strings.SplitN(token, "-", 2)
+ first, parseErr := strconv.ParseInt(parts[0], 10, 64)
+ if parseErr != nil || first < 1 {
+ continue
+ }
+ last := first
+ if len(parts) == 2 {
+ last, parseErr = strconv.ParseInt(parts[1], 10, 64)
+ if parseErr != nil || last < first || last-first > 100000 {
+ continue
+ }
+ }
+ for number := first; number <= last; number++ {
+ articles = append(articles, number)
+ }
+ }
+ }
+ return articles, nil
+}
+
+func (c *Client) Head(number int64) (string, error) {
+ return c.singleArticlePart("HEAD", number, 221)
+}
+
+func (c *Client) Body(number int64) (string, error) {
+ return c.singleArticlePart("BODY", number, 222)
+}
+
+func (c *Client) Stat(number int64) (string, error) {
+ if number < 1 {
+ return "", errors.New("invalid article number")
+ }
+ c.mu.Lock()
+ defer c.mu.Unlock()
+ code, message, err := c.commandUnlocked("STAT " + strconv.FormatInt(number, 10))
+ if err != nil {
+ return "", err
+ }
+ if code != 223 {
+ return "", &ResponseError{Code: code, Message: message}
+ }
+ return message, nil
+}
+
+func (c *Client) Header(field string, first, last int64) ([]HeaderValue, error) {
+ if !validAtom(field) || first < 1 || last < first || last-first > 10000 {
+ return nil, errors.New("invalid HDR request")
+ }
+ c.mu.Lock()
+ defer c.mu.Unlock()
+ lines, err := c.multilineCommand("HDR "+field+" "+strconv.FormatInt(first, 10)+"-"+strconv.FormatInt(last, 10), 225)
+ if err != nil {
+ return nil, err
+ }
+ values := make([]HeaderValue, 0, len(lines))
+ for _, line := range lines {
+ fields := strings.SplitN(line, "\t", 2)
+ if len(fields) != 2 {
+ continue
+ }
+ number, parseErr := strconv.ParseInt(fields[0], 10, 64)
+ if parseErr != nil || number < 1 {
+ continue
+ }
+ values = append(values, HeaderValue{Number: number, Value: fields[1]})
+ }
+ return values, nil
+}
+
+func (c *Client) singleArticlePart(command string, number int64, expected int) (string, error) {
+ if number < 1 {
+ return "", errors.New("invalid article number")
+ }
+ c.mu.Lock()
+ defer c.mu.Unlock()
+ lines, err := c.multilineCommand(command+" "+strconv.FormatInt(number, 10), expected)
+ if err != nil {
+ return "", err
+ }
+ return strings.Join(lines, "\n"), nil
+}
+
+func (c *Client) ListActive() ([]GroupInfo, error) {
+ c.mu.Lock()
+ defer c.mu.Unlock()
+ lines, err := c.multilineCommand("LIST ACTIVE", 215)
+ if err != nil {
+ return nil, err
+ }
+ groups := make([]GroupInfo, 0, len(lines))
+ for _, line := range lines {
+ fields := strings.Fields(line)
+ if len(fields) < 4 || !validAtom(fields[0]) {
+ continue
+ }
+ high, highErr := strconv.ParseInt(fields[1], 10, 64)
+ low, lowErr := strconv.ParseInt(fields[2], 10, 64)
+ if highErr != nil || lowErr != nil || high < 0 || low < 0 {
+ continue
+ }
+ groups = append(groups, GroupInfo{
+ Name: fields[0],
+ Low: low,
+ High: high,
+ EstimatedPost: estimatePopulation(low, high),
+ Posting: fields[3],
+ })
+ }
+ return groups, nil
+}
+
+func (c *Client) SelectGroup(group string) (GroupStatus, error) {
+ if !validAtom(group) {
+ return GroupStatus{}, errors.New("invalid newsgroup name")
+ }
+ c.mu.Lock()
+ defer c.mu.Unlock()
+ return c.selectGroupUnlocked(group)
+}
+
+func (c *Client) selectGroupUnlocked(group string) (GroupStatus, error) {
+ code, message, err := c.commandUnlocked("GROUP " + group)
+ if err != nil {
+ return GroupStatus{}, err
+ }
+ if code != 211 {
+ return GroupStatus{}, &ResponseError{Code: code, Message: message}
+ }
+ fields := strings.Fields(message)
+ if len(fields) < 4 {
+ return GroupStatus{}, errors.New("malformed GROUP response")
+ }
+ count, err := strconv.ParseInt(fields[0], 10, 64)
+ if err != nil {
+ return GroupStatus{}, errors.New("malformed GROUP article count")
+ }
+ low, err := strconv.ParseInt(fields[1], 10, 64)
+ if err != nil {
+ return GroupStatus{}, errors.New("malformed GROUP low article number")
+ }
+ high, err := strconv.ParseInt(fields[2], 10, 64)
+ if err != nil {
+ return GroupStatus{}, errors.New("malformed GROUP high article number")
+ }
+ return GroupStatus{Name: fields[3], Count: count, Low: low, High: high}, nil
+}
+
+func (c *Client) Overview(first, last int64) ([]ArticleHeader, error) {
+ if first < 1 || last < first || last-first > 10_000 {
+ return nil, errors.New("invalid overview range")
+ }
+ c.mu.Lock()
+ defer c.mu.Unlock()
+ return c.overviewUnlocked(first, last)
+}
+
+func (c *Client) overviewUnlocked(first, last int64) ([]ArticleHeader, error) {
+ rangeArg := strconv.FormatInt(first, 10) + "-" + strconv.FormatInt(last, 10)
+ lines, err := c.multilineCommand("OVER "+rangeArg, 224)
+ if code := responseCode(err); code == 500 || code == 501 {
+ lines, err = c.multilineCommand("XOVER "+rangeArg, 224)
+ }
+ if err != nil {
+ return nil, err
+ }
+ headers := make([]ArticleHeader, 0, len(lines))
+ for _, line := range lines {
+ fields := strings.Split(line, "\t")
+ if len(fields) < 5 {
+ continue
+ }
+ number, err := strconv.ParseInt(fields[0], 10, 64)
+ if err != nil {
+ continue
+ }
+ header := ArticleHeader{
+ Number: number,
+ Subject: fields[1],
+ From: fields[2],
+ Date: fields[3],
+ MessageID: fields[4],
+ }
+ if len(fields) > 5 {
+ header.References = fields[5]
+ }
+ if len(fields) > 6 {
+ header.Bytes, _ = strconv.ParseInt(fields[6], 10, 64)
+ }
+ if len(fields) > 7 {
+ header.Lines, _ = strconv.ParseInt(fields[7], 10, 64)
+ }
+ headers = append(headers, header)
+ }
+ return headers, nil
+}
+
+func (c *Client) LatestOverview(group string, limit int64) (GroupStatus, []ArticleHeader, error) {
+ if !validAtom(group) {
+ return GroupStatus{}, nil, errors.New("invalid newsgroup name")
+ }
+ if limit < 1 || limit > 10_000 {
+ return GroupStatus{}, nil, errors.New("overview limit must be between 1 and 10000")
+ }
+ c.mu.Lock()
+ defer c.mu.Unlock()
+ status, err := c.selectGroupUnlocked(group)
+ if err != nil {
+ return GroupStatus{}, nil, err
+ }
+ if status.Count == 0 || status.High < status.Low || status.High < 1 {
+ return status, nil, nil
+ }
+ first := status.High - limit + 1
+ if first < status.Low {
+ first = status.Low
+ }
+ headers, err := c.overviewUnlocked(first, status.High)
+ return status, headers, err
+}
+
+func (c *Client) Article(number int64) (string, error) {
+ if number < 1 {
+ return "", errors.New("invalid article number")
+ }
+ c.mu.Lock()
+ defer c.mu.Unlock()
+ return c.articleUnlocked(number)
+}
+
+func (c *Client) ArticleInGroup(group string, number int64) (string, error) {
+ if !validAtom(group) {
+ return "", errors.New("invalid newsgroup name")
+ }
+ if number < 1 {
+ return "", errors.New("invalid article number")
+ }
+ c.mu.Lock()
+ defer c.mu.Unlock()
+ if _, err := c.selectGroupUnlocked(group); err != nil {
+ return "", err
+ }
+ return c.articleUnlocked(number)
+}
+
+func (c *Client) articleUnlocked(number int64) (string, error) {
+ lines, err := c.multilineCommand("ARTICLE "+strconv.FormatInt(number, 10), 220)
+ if err != nil {
+ return "", err
+ }
+ return strings.Join(lines, "\n"), nil
+}
+
+func (c *Client) Post(article string) error {
+ lines, err := normalizedArticleLines(article)
+ if err != nil {
+ return err
+ }
+ c.mu.Lock()
+ defer c.mu.Unlock()
+ code, message, err := c.commandUnlocked("POST")
+ if err != nil {
+ return err
+ }
+ if code != 340 {
+ return &ResponseError{Code: code, Message: message}
+ }
+ if err := c.setDeadline(); err != nil {
+ return err
+ }
+ for _, line := range lines {
+ if strings.HasPrefix(line, ".") {
+ line = "." + line
+ }
+ if _, err := c.writer.WriteString(line + "\r\n"); err != nil {
+ return fmt.Errorf("write article: %w", err)
+ }
+ }
+ if _, err := c.writer.WriteString(".\r\n"); err != nil {
+ return fmt.Errorf("finish article: %w", err)
+ }
+ if err := c.writer.Flush(); err != nil {
+ return fmt.Errorf("flush article: %w", err)
+ }
+ code, message, err = c.readResponse()
+ if err != nil {
+ return err
+ }
+ if code != 240 {
+ return &ResponseError{Code: code, Message: message}
+ }
+ return nil
+}
+
+// ValidateArticle checks the NNTP/MIME-safe text constraints used by Post.
+// It accepts either LF or CRLF input and does not modify the article.
+func ValidateArticle(article string) error {
+ _, err := normalizedArticleLines(article)
+ return err
+}
+
+func normalizedArticleLines(article string) ([]string, error) {
+ if len(article) == 0 || len(article) > maxPostBytes {
+ return nil, fmt.Errorf("article must contain between 1 and %d bytes", maxPostBytes)
+ }
+ if !utf8.ValidString(article) {
+ return nil, errors.New("article is not valid UTF-8")
+ }
+ if strings.ContainsRune(article, '\x00') {
+ return nil, errors.New("article contains a NUL byte")
+ }
+ normalized := strings.ReplaceAll(article, "\r\n", "\n")
+ normalized = strings.ReplaceAll(normalized, "\r", "\n")
+ lines := strings.Split(normalized, "\n")
+ for _, line := range lines {
+ wireBytes := len(line)
+ if strings.HasPrefix(line, ".") {
+ wireBytes++ // NNTP dot-stuffing adds one byte on the wire.
+ }
+ if wireBytes > maxPostLineBytes {
+ return nil, fmt.Errorf("article contains a line longer than %d octets", maxPostLineBytes)
+ }
+ }
+ return lines, nil
+}
+
+func (c *Client) Close() error {
+ c.mu.Lock()
+ defer c.mu.Unlock()
+ if c.closed {
+ return nil
+ }
+ c.closed = true
+ _ = c.setDeadline()
+ _, _, _ = c.commandUnlocked("QUIT")
+ return c.conn.Close()
+}
+
+func (c *Client) authenticate(username, password string) error {
+ if !validCommandValue(username) || !validCommandValue(password) {
+ return errors.New("credentials contain invalid control characters")
+ }
+ code, message, err := c.command("AUTHINFO USER " + username)
+ if err != nil {
+ return err
+ }
+ if code == 281 {
+ return nil
+ }
+ if code != 381 {
+ return &ResponseError{Code: code, Message: message}
+ }
+ if password == "" {
+ return errors.New("server requires a password")
+ }
+ code, message, err = c.command("AUTHINFO PASS " + password)
+ if err != nil {
+ return err
+ }
+ if code != 281 {
+ return &ResponseError{Code: code, Message: message}
+ }
+ return nil
+}
+
+func (c *Client) authenticateWithConfig(cfg DialConfig) error {
+ if strings.TrimSpace(cfg.SASLMechanism) == "" {
+ return c.authenticate(cfg.Username, cfg.Password)
+ }
+ mechanism := strings.ToUpper(strings.TrimSpace(cfg.SASLMechanism))
+ if mechanism != "PLAIN" {
+ return fmt.Errorf("unsupported NNTP SASL mechanism %q", cfg.SASLMechanism)
+ }
+ capabilities, err := c.capabilitiesUnlocked()
+ if err != nil {
+ return fmt.Errorf("query NNTP capabilities for SASL: %w", err)
+ }
+ if !hasCapability(capabilities, "SASL", "PLAIN") {
+ return errors.New("NNTP server does not advertise SASL PLAIN")
+ }
+ code, message, err := c.commandUnlocked("AUTHINFO SASL PLAIN")
+ if err != nil {
+ return err
+ }
+ if code != 383 {
+ if code == 281 {
+ return nil
+ }
+ return &ResponseError{Code: code, Message: message}
+ }
+ response := base64.StdEncoding.EncodeToString([]byte("\x00" + cfg.Username + "\x00" + cfg.Password))
+ if err := c.writeLineUnlocked(response); err != nil {
+ return err
+ }
+ code, message, err = c.readResponse()
+ if err != nil {
+ return err
+ }
+ if code != 281 {
+ return &ResponseError{Code: code, Message: message}
+ }
+ return nil
+}
+
+func (c *Client) enableCompression() error {
+ capabilities, err := c.capabilitiesUnlocked()
+ if err != nil {
+ return fmt.Errorf("query NNTP capabilities for compression: %w", err)
+ }
+ if !hasCapability(capabilities, "COMPRESS", "DEFLATE") {
+ return errors.New("NNTP server does not advertise COMPRESS DEFLATE")
+ }
+ code, message, err := c.commandUnlocked("COMPRESS DEFLATE")
+ if err != nil {
+ return err
+ }
+ if code != 206 {
+ return &ResponseError{Code: code, Message: message}
+ }
+ compressed := newCompressedConn(c.conn)
+ c.conn = compressed
+ c.reader = bufio.NewReaderSize(compressed, 64*1024)
+ c.writer = bufio.NewWriterSize(compressed, 64*1024)
+ return nil
+}
+
+func hasCapability(lines []string, name string, value string) bool {
+ for _, line := range lines {
+ fields := strings.Fields(line)
+ if len(fields) == 0 || !strings.EqualFold(fields[0], name) {
+ continue
+ }
+ if value == "" {
+ return true
+ }
+ for _, field := range fields[1:] {
+ if strings.EqualFold(field, value) {
+ return true
+ }
+ }
+ }
+ return false
+}
+
+func (c *Client) command(line string) (int, string, error) {
+ c.mu.Lock()
+ defer c.mu.Unlock()
+ return c.commandUnlocked(line)
+}
+
+func (c *Client) commandUnlocked(line string) (int, string, error) {
+ if c.closed && line != "QUIT" {
+ return 0, "", errors.New("NNTP connection is closed")
+ }
+ if !validCommandValue(line) {
+ return 0, "", errors.New("NNTP command contains invalid control characters")
+ }
+ if err := c.setDeadline(); err != nil {
+ return 0, "", err
+ }
+ if err := c.writeLineUnlocked(line); err != nil {
+ return 0, "", err
+ }
+ return c.readResponse()
+}
+
+func (c *Client) writeLineUnlocked(line string) error {
+ if _, err := c.writer.WriteString(line + "\r\n"); err != nil {
+ return fmt.Errorf("write NNTP command: %w", err)
+ }
+ if err := c.writer.Flush(); err != nil {
+ return fmt.Errorf("flush NNTP command: %w", err)
+ }
+ return nil
+}
+
+type compressedConn struct {
+ net.Conn
+ reader io.ReadCloser
+ writer *flate.Writer
+}
+
+func newCompressedConn(conn net.Conn) *compressedConn {
+ writer, _ := flate.NewWriter(conn, flate.DefaultCompression)
+ return &compressedConn{Conn: conn, reader: flate.NewReader(conn), writer: writer}
+}
+
+func (c *compressedConn) Read(p []byte) (int, error) {
+ return c.reader.Read(p)
+}
+
+func (c *compressedConn) Write(p []byte) (int, error) {
+ n, err := c.writer.Write(p)
+ if flushErr := c.writer.Flush(); err == nil {
+ err = flushErr
+ }
+ return n, err
+}
+
+func (c *compressedConn) Close() error {
+ _ = c.writer.Close()
+ _ = c.reader.Close()
+ return c.Conn.Close()
+}
+
+func (c *Client) multilineCommand(command string, expectedCode int) ([]string, error) {
+ code, message, err := c.commandUnlocked(command)
+ if err != nil {
+ return nil, err
+ }
+ if code != expectedCode {
+ return nil, &ResponseError{Code: code, Message: message}
+ }
+ return c.readMultiline()
+}
+
+func (c *Client) readResponse() (int, string, error) {
+ line, err := readLimitedLine(c.reader)
+ if err != nil {
+ return 0, "", err
+ }
+ if len(line) < 3 {
+ return 0, "", errors.New("malformed NNTP response")
+ }
+ code, err := strconv.Atoi(line[:3])
+ if err != nil {
+ return 0, "", errors.New("malformed NNTP response code")
+ }
+ message := ""
+ if len(line) > 4 {
+ message = line[4:]
+ }
+ return code, message, nil
+}
+
+func (c *Client) readMultiline() ([]string, error) {
+ lines := make([]string, 0, 1024)
+ total := 0
+ for len(lines) < maxMultilineLines {
+ line, err := readLimitedLine(c.reader)
+ if err != nil {
+ return nil, err
+ }
+ if line == "." {
+ return lines, nil
+ }
+ if strings.HasPrefix(line, "..") {
+ line = line[1:]
+ }
+ total += len(line)
+ if total > maxMultilineBytes {
+ return nil, errors.New("NNTP multiline response exceeds size limit")
+ }
+ lines = append(lines, line)
+ }
+ return nil, errors.New("NNTP multiline response exceeds line limit")
+}
+
+func (c *Client) setDeadline() error {
+ if err := c.conn.SetDeadline(time.Now().Add(c.timeout)); err != nil {
+ return fmt.Errorf("set NNTP deadline: %w", err)
+ }
+ return nil
+}
+
+func readLimitedLine(reader *bufio.Reader) (string, error) {
+ buffer := make([]byte, 0, 4096)
+ for {
+ fragment, err := reader.ReadSlice('\n')
+ if len(buffer)+len(fragment) > maxResponseLine {
+ return "", errors.New("NNTP response line exceeds size limit")
+ }
+ buffer = append(buffer, fragment...)
+ if errors.Is(err, bufio.ErrBufferFull) {
+ continue
+ }
+ if err != nil {
+ if errors.Is(err, io.EOF) {
+ return "", io.ErrUnexpectedEOF
+ }
+ return "", err
+ }
+ break
+ }
+ line := string(buffer)
+ line = strings.TrimSuffix(line, "\n")
+ line = strings.TrimSuffix(line, "\r")
+ return line, nil
+}
+
+func validateDialConfig(cfg DialConfig) error {
+ if cfg.Host == "" || !validCommandValue(cfg.Host) {
+ return errors.New("invalid NNTP host")
+ }
+ port, err := strconv.Atoi(cfg.Port)
+ if err != nil || port < 1 || port > 65535 {
+ return errors.New("invalid NNTP port")
+ }
+ if cfg.ProxyType != "" && !strings.EqualFold(cfg.ProxyType, "DIRECT") && !strings.EqualFold(cfg.ProxyType, "SOCKS5") {
+ return errors.New("unsupported proxy type")
+ }
+ if strings.EqualFold(cfg.ProxyType, "SOCKS5") {
+ if _, _, err := net.SplitHostPort(cfg.ProxyAddress); err != nil {
+ return fmt.Errorf("invalid SOCKS5 address: %w", err)
+ }
+ }
+ if cfg.Username != "" && !cfg.UseTLS {
+ return errors.New("NNTP authentication requires TLS")
+ }
+ if cfg.StartTLS && !cfg.UseTLS {
+ return errors.New("STARTTLS requires TLS to be enabled")
+ }
+ if cfg.SASLMechanism != "" && cfg.Username == "" {
+ return errors.New("SASL authentication requires a username")
+ }
+ return nil
+}
+
+func makeTLSConfig(cfg DialConfig) *tls.Config {
+ tlsConfig := &tls.Config{
+ MinVersion: tls.VersionTLS12,
+ MaxVersion: tls.VersionTLS13,
+ CipherSuites: append([]uint16(nil), tls12AESGCMSuites...),
+ ServerName: cfg.Host,
+ InsecureSkipVerify: cfg.InsecureSkipVerify,
+ } //nolint:gosec // explicit user opt-in for local or pinned deployments
+ if cfg.tlsConfig != nil {
+ tlsConfig = cfg.tlsConfig.Clone()
+ if tlsConfig.MinVersion < tls.VersionTLS12 {
+ tlsConfig.MinVersion = tls.VersionTLS12
+ }
+ if tlsConfig.MaxVersion < tls.VersionTLS12 || tlsConfig.MaxVersion > tls.VersionTLS13 {
+ tlsConfig.MaxVersion = tls.VersionTLS13
+ }
+ if tlsConfig.CipherSuites == nil {
+ tlsConfig.CipherSuites = append([]uint16(nil), tls12AESGCMSuites...)
+ }
+ if tlsConfig.ServerName == "" {
+ tlsConfig.ServerName = cfg.Host
+ }
+ }
+ return tlsConfig
+}
+
+func estimatePopulation(low, high int64) int64 {
+ if low <= 0 || high < low {
+ return 0
+ }
+ return high - low + 1
+}
+
+func validAtom(value string) bool {
+ return value != "" && validCommandValue(value) && !strings.ContainsAny(value, " ")
+}
+
+func validCommandValue(value string) bool {
+ return !strings.ContainsAny(value, "\x00\r\n")
+}
+
+func responseCode(err error) int {
+ var responseErr *ResponseError
+ if errors.As(err, &responseErr) {
+ return responseErr.Code
+ }
+ return 0
+}