diff options
Diffstat (limited to 'internal/nntp/client.go')
| -rw-r--r-- | internal/nntp/client.go | 983 |
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 +} |
