diff --git a/README.md b/README.md index 3612738..6080c68 100644 --- a/README.md +++ b/README.md @@ -1575,24 +1575,31 @@ reference with all required and optional fields. **Command dispatch table:** -| Command | Required Fields | Optional | Response Status | -| --------- | --------------- | -------- | --------------- | -| `PRIVMSG` | `to`, `body` | `meta` | 200 OK | -| `NOTICE` | `to`, `body` | `meta` | 200 OK | -| `JOIN` | `to` | | 200 OK | -| `PART` | `to` | `body` | 200 OK | -| `NICK` | `body` | | 200 OK | -| `PASS` | `body` | | 200 OK | -| `TOPIC` | `to`, `body` | | 200 OK | -| `MODE` | `to` | | 200 OK | -| `NAMES` | `to` | | 200 OK | -| `LIST` | | | 200 OK | -| `WHOIS` | `to` or `body` | | 200 OK | -| `WHO` | `to` | | 200 OK | -| `LUSERS` | | | 200 OK | -| `OPER` | `body` | | 200 OK | -| `QUIT` | | `body` | 200 OK | -| `PING` | | | 200 OK | +| Command | Required Fields | Optional | Response Status | +| ---------- | --------------- | -------- | --------------- | +| `PRIVMSG` | `to`, `body` | `meta` | 200 OK | +| `NOTICE` | `to`, `body` | `meta` | 200 OK | +| `JOIN` | `to` | | 200 OK | +| `PART` | `to` | `body` | 200 OK | +| `NICK` | `body` | | 200 OK | +| `PASS` | `body` | | 200 OK | +| `TOPIC` | `to`, `body` | | 200 OK | +| `MODE` | `to` | `body` | 200 OK | +| `NAMES` | `to` | | 200 OK | +| `LIST` | | | 200 OK | +| `WHOIS` | `to` or `body` | | 200 OK | +| `WHO` | `to` | | 200 OK | +| `LUSERS` | | | 200 OK | +| `USERHOST` | `body` | | 200 OK | +| `VERSION` | | | 200 OK | +| `ADMIN` | | | 200 OK | +| `INFO` | | | 200 OK | +| `TIME` | | | 200 OK | +| `OPER` | `body` | | 200 OK | +| `KILL` | `body` | | 200 OK | +| `WALLOPS` | `body` | | 200 OK | +| `QUIT` | | `body` | 200 OK | +| `PING` | | | 200 OK | All IRC commands return HTTP 200 OK. IRC-level success and error responses are delivered as **numeric replies** through the message queue (see @@ -1619,7 +1626,11 @@ auth cookies (401), and server errors (500). | 433 | ERR_NICKNAMEINUSE | NICK target is taken | | 442 | ERR_NOTONCHANNEL | Not a member of the target channel | | 461 | ERR_NEEDMOREPARAMS | Missing required fields (to, body) | +| 481 | ERR_NOPRIVILEGES | KILL or WALLOPS by a non-operator | +| 483 | ERR_CANTKILLSERVER | KILL of yourself | | 491 | ERR_NOOPERHOST | Failed OPER authentication | +| 501 | ERR_UMODEUNKNOWNFLAG | User MODE string not accepted | +| 502 | ERR_USERSDONTMATCH | User MODE for a nick other than your own | **IRC numeric success replies (delivered via message queue):** @@ -1630,11 +1641,13 @@ auth cookies (401), and server errors (500). | 003 | RPL_CREATED | Sent on session creation/login | | 004 | RPL_MYINFO | Sent on session creation/login | | 005 | RPL_ISUPPORT | Sent on session creation/login | -| 221 | RPL_UMODEIS | In response to user MODE query | +| 221 | RPL_UMODEIS | In response to user MODE | | 251 | RPL_LUSERCLIENT | On connect or LUSERS command | | 252 | RPL_LUSEROP | On connect or LUSERS command | | 254 | RPL_LUSERCHANNELS | On connect or LUSERS command | | 255 | RPL_LUSERME | On connect or LUSERS command | +| 256–259 | RPL_ADMINME etc. | ADMIN info | +| 302 | RPL_USERHOST | USERHOST reply | | 311 | RPL_WHOISUSER | WHOIS user info | | 312 | RPL_WHOISSERVER | WHOIS server info | | 313 | RPL_WHOISOPERATOR | WHOIS target is oper | @@ -1642,6 +1655,10 @@ auth cookies (401), and server errors (500). | 318 | RPL_ENDOFWHOIS | End of WHOIS list | | 319 | RPL_WHOISCHANNELS | WHOIS channels list | | 338 | RPL_WHOISACTUALLY | WHOIS client IP (oper-only) | +| 351 | RPL_VERSION | VERSION reply | +| 371 | RPL_INFO | INFO line | +| 374 | RPL_ENDOFINFO | End of INFO | +| 391 | RPL_TIME | TIME reply | | 322 | RPL_LIST | Channel in LIST response | | 323 | RPL_LISTEND | End of LIST | | 324 | RPL_CHANNELMODEIS | Channel mode query response | @@ -2384,13 +2401,13 @@ IRC_LISTEN_ADDR= ### Supported Commands -| Category | Commands | -| ---------- | ------------------------------------------------------------------ | -| Connection | `NICK`, `USER`, `PASS`, `QUIT`, `PING`/`PONG`, `CAP` | -| Channels | `JOIN`, `PART`, `MODE`, `TOPIC`, `NAMES`, `LIST`, `KICK`, `INVITE` | -| Messaging | `PRIVMSG`, `NOTICE` | -| Info | `WHO`, `WHOIS`, `LUSERS`, `MOTD`, `AWAY` | -| Operator | `OPER` (requires `NEOIRC_OPER_NAME` and `NEOIRC_OPER_PASSWORD`) | +| Category | Commands | +| ---------- | ---------------------------------------------------------------------------------------- | +| Connection | `NICK`, `USER`, `PASS`, `QUIT`, `PING`/`PONG`, `CAP` | +| Channels | `JOIN`, `PART`, `MODE`, `TOPIC`, `NAMES`, `LIST`, `KICK`, `INVITE` | +| Messaging | `PRIVMSG`, `NOTICE` | +| Info | `WHO`, `WHOIS`, `LUSERS`, `MOTD`, `AWAY`, `USERHOST`, `VERSION`, `ADMIN`, `INFO`, `TIME` | +| Operator | `OPER`, `KILL`, `WALLOPS` (requires `NEOIRC_OPER_NAME` and `NEOIRC_OPER_PASSWORD`) | ### Protocol Details @@ -2406,6 +2423,15 @@ IRC_LISTEN_ADDR= operator status (`@`). - **Channel modes**: `+m` (moderated), `+t` (topic lock), `+o` (operator), `+v` (voice) +- **User modes**: `+o` (operator, set only via `OPER`), `+w` (receives + `WALLOPS`). `MODE` for any nick other than your own is rejected with + `ERR_USERSDONTMATCH` (502), for both queries and changes. Nick comparison is + case-insensitive. +- **KILL**: an operator's `KILL` broadcasts the victim's `QUIT` to its channel + peers, deletes its session, and then sends the victim a `KILL` and + `ERROR :Closing Link` before closing its socket. This applies to victims on + the IRC listener regardless of whether the `KILL` arrived over IRC or the HTTP + API. ### Bridge to HTTP API @@ -2946,6 +2972,10 @@ guess is borne by the server (bcrypt), not the client. from additional devices via `POST /api/v1/login` - [x] **Cookie-based auth** — HttpOnly cookies replace Bearer tokens for all API authentication +- [x] **Tier 3 utility commands** — `USERHOST` (302), `VERSION` (351), `ADMIN` + (256–259), `INFO` (371/374), `TIME` (391), `KILL` (operator-only forced + disconnect), `WALLOPS` (operator-only broadcast to `+w` users) +- [x] **User mode +w** — receive `WALLOPS`, set with `MODE +w` / `-w` ### Future (1.0+) diff --git a/internal/db/export_test.go b/internal/db/export_test.go index 45c0435..f233380 100644 --- a/internal/db/export_test.go +++ b/internal/db/export_test.go @@ -58,3 +58,17 @@ func (database *Database) Close() error { return nil } + +// ExecForTest runs a raw statement, such as a trigger that +// makes a write fail. +func (database *Database) ExecForTest( + ctx context.Context, + query string, +) error { + _, err := database.conn.ExecContext(ctx, query) + if err != nil { + return fmt.Errorf("exec for test: %w", err) + } + + return nil +} diff --git a/internal/db/queries.go b/internal/db/queries.go index d840773..7b437df 100644 --- a/internal/db/queries.go +++ b/internal/db/queries.go @@ -7,6 +7,7 @@ import ( "database/sql" "encoding/hex" "encoding/json" + "errors" "fmt" "strconv" "strings" @@ -2423,3 +2424,172 @@ func (database *Database) SetChannelUserLimit( return nil } + +// SetSessionUserModes sets a session's wallops (+w) and +// oper (+o) flags in one transaction, so that either both +// changes are stored or neither is. A nil flag is left as +// it is. +func (database *Database) SetSessionUserModes( + ctx context.Context, + sessionID int64, + wallops *bool, + oper *bool, +) error { + transaction, err := database.conn.BeginTx(ctx, nil) + if err != nil { + return fmt.Errorf("begin tx: %w", err) + } + + if wallops != nil { + _, err = transaction.ExecContext( + ctx, + `UPDATE sessions SET is_wallops = ? WHERE id = ?`, + boolToInt(*wallops), sessionID, + ) + if err != nil { + _ = transaction.Rollback() + + return fmt.Errorf("set session wallops: %w", err) + } + } + + if oper != nil { + _, err = transaction.ExecContext( + ctx, + `UPDATE sessions SET is_oper = ? WHERE id = ?`, + boolToInt(*oper), sessionID, + ) + if err != nil { + _ = transaction.Rollback() + + return fmt.Errorf("set session oper: %w", err) + } + } + + err = transaction.Commit() + if err != nil { + return fmt.Errorf("commit user modes: %w", err) + } + + return nil +} + +// boolToInt returns the 0 or 1 that the schema stores for a +// boolean. +func boolToInt(value bool) int { + if value { + return 1 + } + + return 0 +} + +// IsSessionWallops returns whether the session has the +// wallops (+w) usermode set. +func (database *Database) IsSessionWallops( + ctx context.Context, + sessionID int64, +) (bool, error) { + var isWallops int + + err := database.conn.QueryRowContext( + ctx, + `SELECT is_wallops FROM sessions WHERE id = ?`, + sessionID, + ).Scan(&isWallops) + if err != nil { + return false, fmt.Errorf( + "check session wallops: %w", err, + ) + } + + return isWallops != 0, nil +} + +// GetWallopsSessionIDs returns all session IDs that have +// the wallops (+w) usermode set. +func (database *Database) GetWallopsSessionIDs( + ctx context.Context, +) ([]int64, error) { + rows, err := database.conn.QueryContext( + ctx, + `SELECT id FROM sessions WHERE is_wallops = 1`, + ) + if err != nil { + return nil, fmt.Errorf( + "get wallops sessions: %w", err, + ) + } + + defer func() { _ = rows.Close() }() + + var ids []int64 + + for rows.Next() { + var sessionID int64 + + err = rows.Scan(&sessionID) + if err != nil { + return nil, fmt.Errorf( + "scan wallops session: %w", err, + ) + } + + ids = append(ids, sessionID) + } + + err = rows.Err() + if err != nil { + return nil, fmt.Errorf( + "iterate wallops sessions: %w", err, + ) + } + + return ids, nil +} + +// UserhostInfo holds the data needed for RPL_USERHOST. +type UserhostInfo struct { + Nick string + Username string + Hostname string + IsOper bool + AwayMessage string +} + +// GetUserhostInfo returns USERHOST info for the given +// nicks. Nicks with no session are left out. +func (database *Database) GetUserhostInfo( + ctx context.Context, + nicks []string, +) ([]UserhostInfo, error) { + results := make([]UserhostInfo, 0, len(nicks)) + + for _, nick := range nicks { + var info UserhostInfo + + err := database.conn.QueryRowContext( + ctx, + `SELECT nick, username, hostname, + is_oper, away_message + FROM sessions WHERE nick = ?`, + nick, + ).Scan( + &info.Nick, &info.Username, &info.Hostname, + &info.IsOper, &info.AwayMessage, + ) + if errors.Is(err, sql.ErrNoRows) { + continue + } + + if err != nil { + return nil, fmt.Errorf( + "userhost lookup %q: %w", nick, err, + ) + } + + results = append(results, info) + } + + return results, nil +} diff --git a/internal/db/queries_test.go b/internal/db/queries_test.go index 9c201d2..d5467a9 100644 --- a/internal/db/queries_test.go +++ b/internal/db/queries_test.go @@ -1491,3 +1491,119 @@ func TestChannelUserLimit(t *testing.T) { t.Fatalf("expected 0, got %d", limit) } } + +// TestSetSessionUserModesIsAtomic makes the is_oper write +// fail after the is_wallops write has run, and checks that +// the wallops change was rolled back. +func TestSetSessionUserModesIsAtomic(t *testing.T) { + t.Parallel() + + database := setupTestDB(t) + ctx := t.Context() + + sessionID, _, _, err := database.CreateSession( + ctx, "alice", "", "", "", + ) + if err != nil { + t.Fatal(err) + } + + err = database.ExecForTest(ctx, + `CREATE TRIGGER reject_oper + BEFORE UPDATE OF is_oper ON sessions + BEGIN SELECT RAISE(ABORT, 'oper write rejected'); + END`, + ) + if err != nil { + t.Fatal(err) + } + + wallops := true + oper := false + + err = database.SetSessionUserModes( + ctx, sessionID, &wallops, &oper, + ) + if err == nil { + t.Fatal("expected the rejected oper write to fail") + } + + gotWallops, err := database.IsSessionWallops( + ctx, sessionID, + ) + if err != nil { + t.Fatal(err) + } + + if gotWallops { + t.Error("wallops was stored although the change failed") + } +} + +// TestSetSessionUserModesAppliesBoth checks that both flags +// are written, and that a nil flag is left as it is. +func TestSetSessionUserModesAppliesBoth(t *testing.T) { + t.Parallel() + + database := setupTestDB(t) + ctx := t.Context() + + sessionID, _, _, err := database.CreateSession( + ctx, "alice", "", "", "", + ) + if err != nil { + t.Fatal(err) + } + + err = database.SetSessionOper(ctx, sessionID, true) + if err != nil { + t.Fatal(err) + } + + wallops := true + oper := false + + err = database.SetSessionUserModes( + ctx, sessionID, &wallops, &oper, + ) + if err != nil { + t.Fatal(err) + } + + gotWallops, err := database.IsSessionWallops( + ctx, sessionID, + ) + if err != nil { + t.Fatal(err) + } + + gotOper, err := database.IsSessionOper(ctx, sessionID) + if err != nil { + t.Fatal(err) + } + + if !gotWallops || gotOper { + t.Errorf( + "want wallops=true oper=false, got %v/%v", + gotWallops, gotOper, + ) + } + + err = database.SetSessionUserModes( + ctx, sessionID, nil, nil, + ) + if err != nil { + t.Fatal(err) + } + + gotWallops, err = database.IsSessionWallops( + ctx, sessionID, + ) + if err != nil { + t.Fatal(err) + } + + if !gotWallops { + t.Error("a nil wallops flag cleared wallops") + } +} diff --git a/internal/db/schema/001_initial.sql b/internal/db/schema/001_initial.sql index e53c48b..1a1b27f 100644 --- a/internal/db/schema/001_initial.sql +++ b/internal/db/schema/001_initial.sql @@ -10,6 +10,7 @@ CREATE TABLE IF NOT EXISTS sessions ( hostname TEXT NOT NULL DEFAULT '', ip TEXT NOT NULL DEFAULT '', is_oper INTEGER NOT NULL DEFAULT 0, + is_wallops INTEGER NOT NULL DEFAULT 0, password_hash TEXT NOT NULL DEFAULT '', signing_key TEXT NOT NULL DEFAULT '', away_message TEXT NOT NULL DEFAULT '', diff --git a/internal/handlers/api.go b/internal/handlers/api.go index 3da43f3..3d98775 100644 --- a/internal/handlers/api.go +++ b/internal/handlers/api.go @@ -1015,10 +1015,10 @@ func (hdlr *Handlers) dispatchCommand( hdlr.handleQuit( writer, request, sessionID, nick, body, ) - case irc.CmdOper: - hdlr.handleOper( + case irc.CmdOper, irc.CmdKill, irc.CmdWallops: + hdlr.dispatchOperCommand( writer, request, - sessionID, clientID, nick, bodyLines, + sessionID, clientID, nick, command, bodyLines, ) case irc.CmdMotd, irc.CmdPing: hdlr.dispatchInfoCommand( @@ -1075,6 +1075,27 @@ func (hdlr *Handlers) dispatchQueryCommand( writer, request, sessionID, clientID, nick, ) + case irc.CmdUserhost: + hdlr.handleUserhost( + writer, request, + sessionID, clientID, nick, bodyLines, + ) + case irc.CmdVersion: + hdlr.handleVersion( + writer, request, sessionID, clientID, nick, + ) + case irc.CmdAdmin: + hdlr.handleAdmin( + writer, request, sessionID, clientID, nick, + ) + case irc.CmdInfo: + hdlr.handleInfo( + writer, request, sessionID, clientID, nick, + ) + case irc.CmdTime: + hdlr.handleTime( + writer, request, sessionID, clientID, nick, + ) default: hdlr.enqueueNumeric( request.Context(), clientID, @@ -1951,15 +1972,11 @@ func (hdlr *Handlers) handleMode( channel := target if !strings.HasPrefix(channel, "#") { - // User mode query — return empty modes. - hdlr.enqueueNumeric( - request.Context(), clientID, - irc.RplUmodeIs, nick, nil, "+", + hdlr.handleUserMode( + writer, request, + sessionID, clientID, nick, target, + bodyLines, ) - hdlr.broker.Notify(sessionID) - hdlr.respondJSON(writer, request, - map[string]string{statusKey: "ok"}, - http.StatusOK) return } diff --git a/internal/handlers/api_test.go b/internal/handlers/api_test.go index 594ef40..4fff74e 100644 --- a/internal/handlers/api_test.go +++ b/internal/handlers/api_test.go @@ -293,6 +293,7 @@ func newTestHandlers( Config: cfg, Database: database, Broker: brk, + Globals: globs, }) hdlr, err := handlers.New(lifecycle, handlers.Params{ //nolint:exhaustruct diff --git a/internal/handlers/utility.go b/internal/handlers/utility.go new file mode 100644 index 0000000..0ae2242 --- /dev/null +++ b/internal/handlers/utility.go @@ -0,0 +1,307 @@ +package handlers + +import ( + "net/http" + "strings" + "time" + + "sneak.berlin/go/neoirc/pkg/irc" +) + +// dispatchOperCommand routes the server operator commands +// OPER, KILL and WALLOPS to their handlers. +func (hdlr *Handlers) dispatchOperCommand( + writer http.ResponseWriter, + request *http.Request, + sessionID, clientID int64, + nick, command string, + bodyLines func() []string, +) { + switch command { + case irc.CmdOper: + hdlr.handleOper( + writer, request, + sessionID, clientID, nick, bodyLines, + ) + case irc.CmdKill: + hdlr.handleKill( + writer, request, + sessionID, clientID, nick, bodyLines, + ) + case irc.CmdWallops: + hdlr.handleWallops( + writer, request, + sessionID, clientID, nick, bodyLines, + ) + } +} + +// handleUserhost handles USERHOST: user@host for each of +// up to five nicks given in the body. +func (hdlr *Handlers) handleUserhost( + writer http.ResponseWriter, + request *http.Request, + sessionID, clientID int64, + nick string, + bodyLines func() []string, +) { + lines := bodyLines() + if len(lines) == 0 { + hdlr.respondIRCError( + writer, request, clientID, sessionID, + irc.ErrNeedMoreParams, nick, + []string{irc.CmdUserhost}, + "Not enough parameters", + ) + + return + } + + reply, err := hdlr.svc.UserhostReply( + request.Context(), lines, hdlr.serverName(), + ) + if hdlr.handleServiceError( + writer, request, clientID, sessionID, nick, err, + ) { + return + } + + hdlr.enqueueNumeric( + request.Context(), clientID, + irc.RplUserHost, nick, nil, reply, + ) + hdlr.broker.Notify(sessionID) + hdlr.respondJSON(writer, request, + map[string]string{statusKey: "ok"}, + http.StatusOK) +} + +// handleVersion handles VERSION. +func (hdlr *Handlers) handleVersion( + writer http.ResponseWriter, + request *http.Request, + sessionID, clientID int64, + nick string, +) { + hdlr.enqueueNumeric( + request.Context(), clientID, irc.RplVersion, nick, + []string{ + hdlr.svc.ServerVersion() + ".", + hdlr.serverName(), + }, + "", + ) + hdlr.broker.Notify(sessionID) + hdlr.respondJSON(writer, request, + map[string]string{statusKey: "ok"}, + http.StatusOK) +} + +// handleAdmin handles ADMIN. +func (hdlr *Handlers) handleAdmin( + writer http.ResponseWriter, + request *http.Request, + sessionID, clientID int64, + nick string, +) { + ctx := request.Context() + srvName := hdlr.serverName() + + hdlr.enqueueNumeric( + ctx, clientID, irc.RplAdminMe, nick, + []string{srvName}, "Administrative info", + ) + hdlr.enqueueNumeric( + ctx, clientID, irc.RplAdminLoc1, nick, nil, + "neoirc server", + ) + hdlr.enqueueNumeric( + ctx, clientID, irc.RplAdminLoc2, nick, nil, + "IRC over HTTP", + ) + hdlr.enqueueNumeric( + ctx, clientID, irc.RplAdminEmail, nick, nil, + "admin@"+srvName, + ) + hdlr.broker.Notify(sessionID) + hdlr.respondJSON(writer, request, + map[string]string{statusKey: "ok"}, + http.StatusOK) +} + +// handleInfo handles INFO. +func (hdlr *Handlers) handleInfo( + writer http.ResponseWriter, + request *http.Request, + sessionID, clientID int64, + nick string, +) { + ctx := request.Context() + + for _, line := range hdlr.svc.InfoLines() { + hdlr.enqueueNumeric( + ctx, clientID, irc.RplInfo, nick, nil, line, + ) + } + + hdlr.enqueueNumeric( + ctx, clientID, irc.RplEndOfInfo, nick, nil, + "End of /INFO list", + ) + hdlr.broker.Notify(sessionID) + hdlr.respondJSON(writer, request, + map[string]string{statusKey: "ok"}, + http.StatusOK) +} + +// handleTime handles TIME: the server's local time. +func (hdlr *Handlers) handleTime( + writer http.ResponseWriter, + request *http.Request, + sessionID, clientID int64, + nick string, +) { + hdlr.enqueueNumeric( + request.Context(), clientID, irc.RplTime, nick, + []string{hdlr.serverName()}, + time.Now().Format(time.RFC1123), + ) + hdlr.broker.Notify(sessionID) + hdlr.respondJSON(writer, request, + map[string]string{statusKey: "ok"}, + http.StatusOK) +} + +// handleKill handles KILL: the body is the target nick and +// an optional reason. +func (hdlr *Handlers) handleKill( + writer http.ResponseWriter, + request *http.Request, + sessionID, clientID int64, + nick string, + bodyLines func() []string, +) { + lines := bodyLines() + + targetNick := "" + if len(lines) > 0 { + targetNick = strings.TrimSpace(lines[0]) + } + + if targetNick == "" { + hdlr.respondIRCError( + writer, request, clientID, sessionID, + irc.ErrNeedMoreParams, nick, + []string{irc.CmdKill}, + "Not enough parameters", + ) + + return + } + + reason := "KILLed" + if len(lines) > 1 { + reason = lines[1] + } + + err := hdlr.svc.KillUser( + request.Context(), sessionID, nick, + targetNick, reason, + ) + if hdlr.handleServiceError( + writer, request, clientID, sessionID, nick, err, + ) { + return + } + + hdlr.respondJSON(writer, request, + map[string]string{statusKey: "ok"}, + http.StatusOK) +} + +// handleWallops handles WALLOPS: the body is the message. +func (hdlr *Handlers) handleWallops( + writer http.ResponseWriter, + request *http.Request, + sessionID, clientID int64, + nick string, + bodyLines func() []string, +) { + lines := bodyLines() + if len(lines) == 0 { + hdlr.respondIRCError( + writer, request, clientID, sessionID, + irc.ErrNeedMoreParams, nick, + []string{irc.CmdWallops}, + "Not enough parameters", + ) + + return + } + + err := hdlr.svc.SendWallops( + request.Context(), sessionID, nick, + strings.Join(lines, " "), + ) + if hdlr.handleServiceError( + writer, request, clientID, sessionID, nick, err, + ) { + return + } + + hdlr.respondJSON(writer, request, + map[string]string{statusKey: "ok"}, + http.StatusOK) +} + +// handleUserMode handles MODE for a nick: a query when the +// body is empty, otherwise a change by the mode string in +// the body. Only the caller's own nick, in any letter case, +// is allowed. +func (hdlr *Handlers) handleUserMode( + writer http.ResponseWriter, + request *http.Request, + sessionID, clientID int64, + nick, target string, + bodyLines func() []string, +) { + if !strings.EqualFold(target, nick) { + hdlr.respondIRCError( + writer, request, clientID, sessionID, + irc.ErrUsersDoNotMatch, nick, nil, + "Can't change mode for other users", + ) + + return + } + + ctx := request.Context() + + var ( + modes string + err error + ) + + lines := bodyLines() + if len(lines) > 0 { + modes, err = hdlr.svc.ApplyUserMode( + ctx, sessionID, lines[0], + ) + } else { + modes, err = hdlr.svc.QueryUserMode(ctx, sessionID) + } + + if hdlr.handleServiceError( + writer, request, clientID, sessionID, nick, err, + ) { + return + } + + hdlr.enqueueNumeric( + ctx, clientID, irc.RplUmodeIs, nick, nil, modes, + ) + hdlr.broker.Notify(sessionID) + hdlr.respondJSON(writer, request, + map[string]string{statusKey: "ok"}, + http.StatusOK) +} diff --git a/internal/handlers/utility_test.go b/internal/handlers/utility_test.go new file mode 100644 index 0000000..cb76075 --- /dev/null +++ b/internal/handlers/utility_test.go @@ -0,0 +1,324 @@ +//nolint:paralleltest // global viper, as in api_test.go +package handlers_test + +import ( + "net/http" + "strings" + "testing" + + "sneak.berlin/go/neoirc/pkg/irc" +) + +// Each test starts one server and reuses it for all its +// checks, and makes as few requests as it can: servers, +// sessions and requests are what makes this package's tests +// slow. + +// sendAndPoll sends cmd from the session token and returns +// the messages after lastID that the session has been sent, +// and the ID of the last of them. +func sendAndPoll( + tserver *testServer, + token string, + lastID int64, + cmd map[string]any, +) ([]map[string]any, int64) { + tserver.t.Helper() + + tserver.sendCommand(token, cmd) + + return tserver.pollMessages(token, lastID) +} + +// newOperSession creates a session for nick, makes it a +// server operator, and returns its token and the ID of the +// last message it has been sent. +func newOperSession( + tserver *testServer, nick string, +) (string, int64) { + tserver.t.Helper() + + token := tserver.createSession(nick) + + _, lastID := sendAndPoll(tserver, token, 0, map[string]any{ + commandKey: irc.CmdOper, + bodyKey: []string{testOperName, testOperPassword}, + }) + + return token, lastID +} + +// numericBody returns the first body line of the message in +// msgs with the given numeric, and fails the test if there +// is none. +func numericBody( + t *testing.T, msgs []map[string]any, numeric string, +) string { + t.Helper() + + msg := findNumericWithParams(msgs, numeric) + if msg == nil { + t.Fatalf("expected numeric %s, got %v", numeric, msgs) + } + + lines, _ := msg[bodyKey].([]any) + if len(lines) == 0 { + return "" + } + + line, _ := lines[0].(string) + + return line +} + +func TestUserhost(t *testing.T) { + tserver := newTestServer(t) + + token := tserver.createSession("alice") + tserver.createSession("bob") + + _, lastID := tserver.pollMessages(token, 0) + + msgs, lastID := sendAndPoll(tserver, token, lastID, map[string]any{ + commandKey: irc.CmdUserhost, + bodyKey: []string{"alice", "bob"}, + }) + + body := numericBody(t, msgs, "302") + if !strings.HasPrefix(body, "alice=+alice@") || + !strings.Contains(body, " bob=+bob@") { + t.Errorf("expected alice and bob, got %q", body) + } + + msgs, _ = sendAndPoll(tserver, token, lastID, map[string]any{ + commandKey: irc.CmdUserhost, + }) + if !findNumeric(msgs, "461") { + t.Errorf("expected ERR_NEEDMOREPARAMS (461), got %v", msgs) + } +} + +func TestVersionAdminInfoTime(t *testing.T) { + tserver := newTestServer(t) + + token := tserver.createSession("frank") + _, lastID := tserver.pollMessages(token, 0) + + msgs, lastID := sendAndPoll(tserver, token, lastID, map[string]any{ + commandKey: irc.CmdVersion, + }) + + params := getNumericParams(findNumericWithParams(msgs, "351")) + if len(params) == 0 || params[0] != "neoirc-test-test." { + t.Errorf("expected RPL_VERSION neoirc-test-test., got %v", msgs) + } + + for _, check := range []struct { + command string + numerics []string + }{ + {irc.CmdAdmin, []string{"256", "257", "258", "259"}}, + {irc.CmdInfo, []string{"371", "374"}}, + {irc.CmdTime, []string{"391"}}, + } { + msgs, lastID = sendAndPoll(tserver, token, lastID, map[string]any{ + commandKey: check.command, + }) + + for _, numeric := range check.numerics { + if !findNumeric(msgs, numeric) { + t.Errorf( + "%s: expected %s, got %v", + check.command, numeric, msgs, + ) + } + } + } +} + +func TestKill(t *testing.T) { + tserver := newTestServerWithOper(t) + + victimToken := tserver.createSession("victim") + observerToken := tserver.createSession("observer") + notOperToken := tserver.createSession("notoper") + operToken, operLastID := newOperSession(tserver, "killer") + + for _, token := range []string{victimToken, observerToken} { + tserver.sendCommand(token, map[string]any{ + commandKey: joinCmd, + toKey: "#killtest", + }) + } + + _, notOperLastID := tserver.pollMessages(notOperToken, 0) + + msgs, _ := sendAndPoll(tserver, notOperToken, notOperLastID, map[string]any{ + commandKey: irc.CmdKill, + bodyKey: []string{"victim"}, + }) + if !findNumeric(msgs, "481") { + t.Errorf("expected ERR_NOPRIVILEGES (481), got %v", msgs) + } + + for _, check := range []struct { + numeric string + body []string + }{ + {"461", []string{}}, + {"401", []string{"ghost"}}, + {"483", []string{"killer"}}, + } { + msgs, operLastID = sendAndPoll(tserver, operToken, operLastID, map[string]any{ + commandKey: irc.CmdKill, + bodyKey: check.body, + }) + if !findNumeric(msgs, check.numeric) { + t.Errorf( + "KILL %v: expected %s, got %v", + check.body, check.numeric, msgs, + ) + } + } + + _, observerLastID := tserver.pollMessages(observerToken, 0) + + status, result := tserver.sendCommand(operToken, map[string]any{ + commandKey: irc.CmdKill, + bodyKey: []string{"victim", "go away"}, + }) + if status != http.StatusOK { + t.Fatalf("expected 200, got %d: %v", status, result) + } + + msgs, _ = tserver.pollMessages(observerToken, observerLastID) + if !findMessage(msgs, irc.CmdQuit, "victim") { + t.Errorf("expected the observer to see QUIT, got %v", msgs) + } + + status, _ = tserver.getState(victimToken) + if status != http.StatusUnauthorized { + t.Errorf("expected the victim's session gone, got %d", status) + } +} + +func TestWallops(t *testing.T) { + tserver := newTestServerWithOper(t) + + receiverToken := tserver.createSession("receiver") + plainToken := tserver.createSession("plain") + operToken, operLastID := newOperSession(tserver, "walloper") + + _, receiverLastID := sendAndPoll(tserver, receiverToken, 0, map[string]any{ + commandKey: modeCmd, + toKey: "receiver", + bodyKey: []string{"+w"}, + }) + _, plainLastID := tserver.pollMessages(plainToken, 0) + + msgs, plainLastID := sendAndPoll(tserver, plainToken, plainLastID, map[string]any{ + commandKey: irc.CmdWallops, + bodyKey: []string{"not allowed"}, + }) + if !findNumeric(msgs, "481") { + t.Errorf("expected ERR_NOPRIVILEGES (481), got %v", msgs) + } + + msgs, _ = sendAndPoll(tserver, operToken, operLastID, map[string]any{ + commandKey: irc.CmdWallops, + }) + if !findNumeric(msgs, "461") { + t.Errorf("expected ERR_NEEDMOREPARAMS (461), got %v", msgs) + } + + tserver.sendCommand(operToken, map[string]any{ + commandKey: irc.CmdWallops, + bodyKey: []string{"server going down"}, + }) + + msgs, _ = tserver.pollMessages(receiverToken, receiverLastID) + if !findMessage(msgs, irc.CmdWallops, "walloper") { + t.Errorf("expected WALLOPS for the +w user, got %v", msgs) + } + + msgs, _ = tserver.pollMessages(plainToken, plainLastID) + if findMessage(msgs, irc.CmdWallops, "walloper") { + t.Errorf("WALLOPS reached a user without +w: %v", msgs) + } +} + +func TestUserMode(t *testing.T) { + const ( + nick = "alice" + other = "other" + ) + + tserver := newTestServerWithOper(t) + + token := tserver.createSession(nick) + otherToken := tserver.createSession(other) + operToken, operLastID := newOperSession(tserver, "deoper") + + // mode sends MODE for target from token, with modeStr + // unless it is empty, and returns the session's next + // messages. + mode := func( + token string, lastID int64, target, modeStr string, + ) ([]map[string]any, int64) { + cmd := map[string]any{commandKey: modeCmd, toKey: target} + if modeStr != "" { + cmd[bodyKey] = []string{modeStr} + } + + return sendAndPoll(tserver, token, lastID, cmd) + } + + _, otherLastID := mode(otherToken, 0, other, "+w") + _, lastID := tserver.pollMessages(token, 0) + + for _, check := range []struct { + target, modeStr, numeric, modes string + }{ + {nick, "+w", "221", "+w"}, + {nick, "", "221", "+w"}, + {strings.ToUpper(nick), "-w", "221", "+"}, + {nick, "+z", "501", ""}, + {other, "", "502", ""}, + {other, "-w", "502", ""}, + } { + var msgs []map[string]any + + msgs, lastID = mode(token, lastID, check.target, check.modeStr) + + if check.numeric != "221" { + if !findNumeric(msgs, check.numeric) || + findNumeric(msgs, "221") { + t.Errorf( + "MODE %s %s: expected only %s, got %v", + check.target, check.modeStr, + check.numeric, msgs, + ) + } + + continue + } + + got := numericBody(t, msgs, "221") + if got != check.modes { + t.Errorf( + "MODE %s %s: expected %q, got %q", + check.target, check.modeStr, check.modes, got, + ) + } + } + + msgs, _ := mode(otherToken, otherLastID, other, "") + if got := numericBody(t, msgs, "221"); got != "+w" { + t.Errorf("%s's modes changed to %q", other, got) + } + + msgs, _ = mode(operToken, operLastID, "deoper", "-o") + if got := numericBody(t, msgs, "221"); got != "+" { + t.Errorf("after -o: expected +, got %q", got) + } +} diff --git a/internal/ircserver/commands.go b/internal/ircserver/commands.go index e6c3356..99785ae 100644 --- a/internal/ircserver/commands.go +++ b/internal/ircserver/commands.go @@ -13,7 +13,7 @@ import ( ) // sendIRCError maps a service.IRCError to an IRC numeric -// reply on the wire. +// reply on the wire, and logs any other error. func (c *Conn) sendIRCError(err error) { var ircErr *service.IRCError if errors.As(err, &ircErr) { @@ -21,7 +21,11 @@ func (c *Conn) sendIRCError(err error) { args = append(args, ircErr.Params...) args = append(args, ircErr.Message) c.sendNumeric(ircErr.Code, args...) + + return } + + c.log.Error("command failed", "error", err) } // handleCAP silently acknowledges CAP negotiation. @@ -345,7 +349,10 @@ func (c *Conn) handleQuit(msg *Message) { c.send("ERROR :Closing Link: " + c.hostname + " (Quit: " + reason + ")") + + c.mu.Lock() c.closed = true + c.mu.Unlock() } // handleTopic gets or sets a channel topic via the shared @@ -427,7 +434,7 @@ func (c *Conn) handleMode( if strings.HasPrefix(target, "#") { c.handleChannelMode(ctx, msg) } else { - c.handleUserMode(msg) + c.handleUserMode(ctx, msg) } } @@ -686,11 +693,14 @@ func (c *Conn) applyChannelModes( } } -// handleUserMode handles MODE for users. -func (c *Conn) handleUserMode(msg *Message) { - target := msg.Params[0] - - if !strings.EqualFold(target, c.nick) { +// handleUserMode handles MODE for a nick: a query without a +// mode string, otherwise a change. Only the client's own +// nick, in any letter case, is allowed. +func (c *Conn) handleUserMode( + ctx context.Context, + msg *Message, +) { + if !strings.EqualFold(msg.Params[0], c.currentNick()) { c.sendNumeric( irc.ErrUsersDoNotMatch, "Can't change mode for other users", @@ -699,8 +709,26 @@ func (c *Conn) handleUserMode(msg *Message) { return } - // We don't support user modes beyond the basics. - c.sendNumeric(irc.RplUmodeIs, "+") + var ( + modes string + err error + ) + + if len(msg.Params) == 1 { + modes, err = c.svc.QueryUserMode(ctx, c.sessionID) + } else { + modes, err = c.svc.ApplyUserMode( + ctx, c.sessionID, msg.Params[1], + ) + } + + if err != nil { + c.sendIRCError(err) + + return + } + + c.sendNumeric(irc.RplUmodeIs, modes) } // handleNames replies with channel member list. @@ -1248,47 +1276,114 @@ func (c *Conn) handleInvite( ) } -// handleUserhost replies with USERHOST info. +// handleUserhost replies with user@host for each of up to +// five nicks. func (c *Conn) handleUserhost( ctx context.Context, msg *Message, ) { if len(msg.Params) < 1 { + c.sendNumeric( + irc.ErrNeedMoreParams, + irc.CmdUserhost, "Not enough parameters", + ) + return } - replies := make([]string, 0, len(msg.Params)) + reply, err := c.svc.UserhostReply( + ctx, msg.Params, c.serverSfx, + ) + if err != nil { + c.sendIRCError(err) - for _, nick := range msg.Params { - sid, err := c.database.GetSessionByNick(ctx, nick) - if err != nil { - continue - } - - hostInfo, _ := c.database.GetSessionHostInfo( - ctx, sid, - ) - - host := "*" - if hostInfo != nil { - host = hostInfo.Hostname - } - - isOper, _ := c.database.IsSessionOper(ctx, sid) - - operStar := "" - if isOper { - operStar = "*" - } - - replies = append( - replies, - nick+operStar+"=+"+nick+"@"+host, - ) + return } + c.sendNumeric(irc.RplUserHost, reply) +} + +// handleVersion replies with the server version. +func (c *Conn) handleVersion() { c.sendNumeric( - irc.RplUserHost, - strings.Join(replies, " "), + irc.RplVersion, + c.svc.ServerVersion()+".", c.serverSfx, "", ) } + +// handleAdmin replies with the server's admin info. +func (c *Conn) handleAdmin() { + c.sendNumeric( + irc.RplAdminMe, c.serverSfx, "Administrative info", + ) + c.sendNumeric(irc.RplAdminLoc1, "neoirc server") + c.sendNumeric(irc.RplAdminLoc2, "IRC over HTTP") + c.sendNumeric(irc.RplAdminEmail, "admin@"+c.serverSfx) +} + +// handleInfo replies with the server's software info. +func (c *Conn) handleInfo() { + for _, line := range c.svc.InfoLines() { + c.sendNumeric(irc.RplInfo, line) + } + + c.sendNumeric(irc.RplEndOfInfo, "End of /INFO list") +} + +// handleTime replies with the server's local time. +func (c *Conn) handleTime() { + c.sendNumeric( + irc.RplTime, + c.serverSfx, time.Now().Format(time.RFC1123), + ) +} + +// handleKill handles KILL []. +func (c *Conn) handleKill( + ctx context.Context, + msg *Message, +) { + if len(msg.Params) < 1 { + c.sendNumeric( + irc.ErrNeedMoreParams, + irc.CmdKill, "Not enough parameters", + ) + + return + } + + reason := "KILLed" + if len(msg.Params) > 1 { + reason = msg.Params[1] + } + + err := c.svc.KillUser( + ctx, c.sessionID, c.currentNick(), + msg.Params[0], reason, + ) + if err != nil { + c.sendIRCError(err) + } +} + +// handleWallops handles WALLOPS . +func (c *Conn) handleWallops( + ctx context.Context, + msg *Message, +) { + if len(msg.Params) < 1 { + c.sendNumeric( + irc.ErrNeedMoreParams, + irc.CmdWallops, "Not enough parameters", + ) + + return + } + + err := c.svc.SendWallops( + ctx, c.sessionID, c.currentNick(), msg.Params[0], + ) + if err != nil { + c.sendIRCError(err) + } +} diff --git a/internal/ircserver/conn.go b/internal/ircserver/conn.go index fdd6b35..74e4dd1 100644 --- a/internal/ircserver/conn.go +++ b/internal/ircserver/conn.go @@ -22,6 +22,7 @@ const ( maxLineLen = 512 readTimeout = 5 * time.Minute writeTimeout = 30 * time.Second + killWriteWindow = 2 * time.Second dnsTimeout = 3 * time.Second pollInterval = 100 * time.Millisecond pingInterval = 90 * time.Second @@ -46,6 +47,11 @@ type Conn struct { serverSfx string commands map[string]cmdHandler + // writeMu serializes writes to conn, which come from + // serve(), from the relay goroutine and, for KILL, from + // Disconnect. + writeMu sync.Mutex + mu sync.Mutex nick string username string @@ -62,6 +68,7 @@ type Conn struct { lastQueueID int64 closed bool + killed bool } func newConn( @@ -97,6 +104,57 @@ func newConn( return conn } +// Disconnect ends the connection for an operator KILL from +// either transport: the victim is sent KILL and ERROR, then +// its socket is closed, which ends serve() and with it the +// relay goroutine. It is called from the killer's goroutine +// and does not wait for the writes, so a victim that has +// stopped reading cannot stall the killer. +func (c *Conn) Disconnect(reason string) { + c.mu.Lock() + + if c.closed { + c.mu.Unlock() + + return + } + + c.closed = true + c.killed = true + nick := c.nick + host := c.hostname + c.mu.Unlock() + + if nick == "" { + nick = "*" + } + + go func() { + c.notifyKilledAndClose(nick, host, reason) + }() +} + +// notifyKilledAndClose sends a killed victim the KILL and +// ERROR lines, each bounded by killWriteWindow, and then +// closes its socket. +func (c *Conn) notifyKilledAndClose( + nick, host, reason string, +) { + defer func() { _ = c.conn.Close() }() + + c.sendWithin( + killWriteWindow, + FormatMessage( + c.serverSfx, irc.CmdKill, nick, reason, + ), + ) + c.sendWithin( + killWriteWindow, + "ERROR :Closing Link: "+host+ + " ("+reason+")", + ) +} + // buildCommandMap returns a map from IRC command strings // to handler functions. func (c *Conn) buildCommandMap() map[string]cmdHandler { @@ -129,7 +187,13 @@ func (c *Conn) buildCommandMap() map[string]cmdHandler { "CAP": func(_ context.Context, msg *Message) { c.handleCAP(msg) }, - "USERHOST": c.handleUserhost, + irc.CmdUserhost: c.handleUserhost, + irc.CmdVersion: func(context.Context, *Message) { c.handleVersion() }, + irc.CmdAdmin: func(context.Context, *Message) { c.handleAdmin() }, + irc.CmdInfo: func(context.Context, *Message) { c.handleInfo() }, + irc.CmdTime: func(context.Context, *Message) { c.handleTime() }, + irc.CmdKill: c.handleKill, + irc.CmdWallops: c.handleWallops, } } @@ -180,7 +244,7 @@ func (c *Conn) serve(ctx context.Context) { c.handleMessage(ctx, msg) - if c.closed { + if c.isClosed() { return } } @@ -189,24 +253,60 @@ func (c *Conn) serve(ctx context.Context) { func (c *Conn) cleanup(ctx context.Context) { c.mu.Lock() wasRegistered := c.registered + wasKilled := c.killed sessID := c.sessionID nick := c.nick c.closed = true c.mu.Unlock() if wasRegistered && sessID > 0 { - c.svc.BroadcastQuit( - ctx, sessID, nick, "Connection closed", - ) + c.svc.UnregisterWireConn(sessID, c) + + // KILL has already sent the QUIT and deleted the + // session. + if !wasKilled { + c.svc.BroadcastQuit( + ctx, sessID, nick, "Connection closed", + ) + } } c.conn.Close() //nolint:errcheck,gosec } +// isClosed reports whether QUIT or KILL has ended the +// connection. +func (c *Conn) isClosed() bool { + c.mu.Lock() + defer c.mu.Unlock() + + return c.closed +} + +// currentNick returns c.nick, read under c.mu. +func (c *Conn) currentNick() string { + c.mu.Lock() + defer c.mu.Unlock() + + return c.nick +} + // send writes a formatted IRC line to the connection. func (c *Conn) send(line string) { + c.sendWithin(writeTimeout, line) +} + +// sendWithin writes a formatted IRC line to the connection, +// giving up after timeout. +func (c *Conn) sendWithin( + timeout time.Duration, + line string, +) { + c.writeMu.Lock() + defer c.writeMu.Unlock() + _ = c.conn.SetWriteDeadline( - time.Now().Add(writeTimeout), + time.Now().Add(timeout), ) _, _ = fmt.Fprintf(c.conn, "%s\r\n", line) @@ -387,7 +487,10 @@ func (c *Conn) completeRegistration(ctx context.Context) { "failed to create session", "error", err, ) c.send("ERROR :Internal server error") + + c.mu.Lock() c.closed = true + c.mu.Unlock() return } @@ -398,6 +501,9 @@ func (c *Conn) completeRegistration(ctx context.Context) { c.registered = true c.mu.Unlock() + // So that KILL, from either transport, can close it. + c.svc.RegisterWireConn(sessionID, c) + // If PASS was provided before registration, set the // session password. if c.passWord != "" && len(c.passWord) >= minPasswordLen { diff --git a/internal/ircserver/conn_test.go b/internal/ircserver/conn_test.go new file mode 100644 index 0000000..75601a1 --- /dev/null +++ b/internal/ircserver/conn_test.go @@ -0,0 +1,100 @@ +package ircserver_test + +import ( + "bufio" + "log/slog" + "net" + "os" + "strings" + "testing" + "time" + + "sneak.berlin/go/neoirc/internal/config" + "sneak.berlin/go/neoirc/internal/ircserver" +) + +// newKillVictim returns a Conn for the nick victim whose +// socket is one end of a net.Pipe, and the other end, from +// which the test reads what the victim is sent. A write to +// the pipe blocks until the other end reads. +func newKillVictim(t *testing.T) (*ircserver.Conn, net.Conn) { + t.Helper() + + serverSide, clientSide := net.Pipe() + + t.Cleanup(func() { _ = clientSide.Close() }) + + log := slog.New(slog.NewTextHandler( + os.Stderr, + &slog.HandlerOptions{Level: slog.LevelError}, + )) + cfg := &config.Config{ServerName: testServerName} + + return ircserver.NewTestConn( + log, cfg, serverSide, "victim", + ), clientSide +} + +// TestDisconnectDoesNotBlockOnUnresponsiveVictim checks +// that Disconnect returns although the victim never reads, +// so a KILL cannot stall the operator who sent it. +func TestDisconnectDoesNotBlockOnUnresponsiveVictim( + t *testing.T, +) { + t.Parallel() + + victim, _ := newKillVictim(t) + returned := make(chan struct{}) + + go func() { + victim.Disconnect("killed by oper") + close(returned) + }() + + select { + case <-returned: + case <-time.After(testTimeout): + t.Fatal("Disconnect waited for the victim to read") + } +} + +// TestDisconnectNotifiesAndClosesVictim checks that a +// victim that reads is sent KILL and ERROR, and that its +// connection is then closed. +func TestDisconnectNotifiesAndClosesVictim(t *testing.T) { + t.Parallel() + + victim, clientSide := newKillVictim(t) + received := make(chan []string, 1) + + go func() { + var lines []string + + scanner := bufio.NewScanner(clientSide) + for scanner.Scan() { + lines = append(lines, scanner.Text()) + } + + received <- lines + }() + + victim.Disconnect("killed by oper") + + var lines []string + + select { + case lines = <-received: + case <-time.After(testTimeout): + t.Fatal("the victim's connection was not closed") + } + + joined := strings.Join(lines, "\n") + + if !strings.Contains(joined, "KILL victim :killed by oper") { + t.Errorf("missing KILL line, got: %q", joined) + } + + if !strings.Contains(joined, "ERROR :Closing Link:") { + t.Errorf("missing ERROR line, got: %q", joined) + } +} diff --git a/internal/ircserver/export_test.go b/internal/ircserver/export_test.go index 84e870b..2b8e08c 100644 --- a/internal/ircserver/export_test.go +++ b/internal/ircserver/export_test.go @@ -4,10 +4,12 @@ import ( "context" "log/slog" "net" + "time" "sneak.berlin/go/neoirc/internal/broker" "sneak.berlin/go/neoirc/internal/config" "sneak.berlin/go/neoirc/internal/db" + "sneak.berlin/go/neoirc/internal/globals" "sneak.berlin/go/neoirc/internal/service" ) @@ -19,8 +21,14 @@ func NewTestServer( database *db.Database, brk *broker.Broker, ) *Server { + globs := &globals.Globals{ + Appname: "neoirc", + Version: "test", + StartTime: time.Now(), + } + svc := service.NewTestService( - database, brk, cfg, log, + database, brk, cfg, globs, log, ) return &Server{ //nolint:exhaustruct @@ -47,3 +55,20 @@ func (s *Server) Stop() { func (s *Server) Listener() net.Listener { return s.listener } + +// NewTestConn returns a Conn for nick on tcpConn, without a +// server, database or service behind it. +func NewTestConn( + log *slog.Logger, + cfg *config.Config, + tcpConn net.Conn, + nick string, +) *Conn { + conn := newConn( + context.Background(), tcpConn, log, + nil, nil, cfg, nil, + ) + conn.nick = nick + + return conn +} diff --git a/internal/ircserver/integration_test.go b/internal/ircserver/integration_test.go index 57cd413..adb69f4 100644 --- a/internal/ircserver/integration_test.go +++ b/internal/ircserver/integration_test.go @@ -4,6 +4,8 @@ import ( "strings" "testing" "time" + + "sneak.berlin/go/neoirc/internal/config" ) // TestIntegrationTwoClients is a comprehensive integration @@ -762,6 +764,371 @@ func TestIntegrationTwoClients(t *testing.T) { ) } +// ── Tier 3 Utility Command Integration Tests ────────── + +// TestIntegrationUserhost verifies the USERHOST command +// returns user@host info for connected nicks. +func TestIntegrationUserhost(t *testing.T) { + t.Parallel() + + env := newTestEnv(t) + + alice := env.dial(t) + alice.register("alice") + + bob := env.dial(t) + bob.register("bob") + + bob.send("AWAY :lunch") + bob.readUntil(func(l string) bool { + return strings.Contains(l, " 306 ") + }) + + for _, step := range []struct { + line, want string + }{ + {"USERHOST alice", " 302 alice alice=+alice@"}, + {"USERHOST alice bob", " :alice=+alice@"}, + {"USERHOST alice bob", " bob=-bob@"}, + {"USERHOST nobody", " 302 alice :"}, + {"USERHOST", " 461 alice USERHOST :"}, + } { + alice.send(step.line) + + reply := alice.readUntil(func(l string) bool { + return strings.Contains(l, " 302 ") || + strings.Contains(l, " 461 ") + }) + assertContains(t, reply, step.want, step.line) + } +} + +// TestIntegrationVersion verifies the VERSION command +// returns the server version string. +func TestIntegrationVersion(t *testing.T) { + t.Parallel() + + env := newTestEnv(t) + + alice := env.dial(t) + alice.register("alice") + + alice.send("VERSION") + + aliceReply := alice.readUntil(func(l string) bool { + return strings.Contains(l, " 351 ") + }) + assertContains( + t, aliceReply, " 351 ", + "RPL_VERSION", + ) + assertContains( + t, aliceReply, "neoirc", + "VERSION reply contains server name", + ) +} + +// TestIntegrationAdmin verifies the ADMIN command returns +// server admin info (256–259 numerics). +func TestIntegrationAdmin(t *testing.T) { + t.Parallel() + + env := newTestEnv(t) + + alice := env.dial(t) + alice.register("alice") + + alice.send("ADMIN") + + aliceReply := alice.readUntil(func(l string) bool { + return strings.Contains(l, " 259 ") + }) + assertContains( + t, aliceReply, " 256 ", + "RPL_ADMINME", + ) + assertContains( + t, aliceReply, " 257 ", + "RPL_ADMINLOC1", + ) + assertContains( + t, aliceReply, " 258 ", + "RPL_ADMINLOC2", + ) + assertContains( + t, aliceReply, " 259 ", + "RPL_ADMINEMAIL", + ) +} + +// TestIntegrationInfo verifies the INFO command returns +// server information (371/374 numerics). +func TestIntegrationInfo(t *testing.T) { + t.Parallel() + + env := newTestEnv(t) + + alice := env.dial(t) + alice.register("alice") + + alice.send("INFO") + + aliceReply := alice.readUntil(func(l string) bool { + return strings.Contains(l, " 374 ") + }) + assertContains( + t, aliceReply, " 371 ", + "RPL_INFO", + ) + assertContains( + t, aliceReply, " 374 ", + "RPL_ENDOFINFO", + ) + assertContains( + t, aliceReply, "neoirc", + "INFO reply mentions server name", + ) +} + +// TestIntegrationTime verifies the TIME command returns +// the server time (391 numeric). +func TestIntegrationTime(t *testing.T) { + t.Parallel() + + env := newTestEnv(t) + + alice := env.dial(t) + alice.register("alice") + + alice.send("TIME") + + aliceReply := alice.readUntil(func(l string) bool { + return strings.Contains(l, " 391 ") + }) + assertContains( + t, aliceReply, " 391 ", + "RPL_TIME", + ) + assertContains( + t, aliceReply, testServerName, + "TIME reply includes server name", + ) +} + +// becomeOper sends OPER with newTestEnvWithOper's +// credentials and waits for RPL_YOUREOPER. +func (tc *testClient) becomeOper() { + tc.t.Helper() + + tc.send("OPER testoper testpass") + tc.readUntil(func(l string) bool { + return strings.Contains(l, " 381 ") + }) +} + +// TestIntegrationKillRefused covers the KILL errors: not an +// operator, no such nick, and killing yourself. +func TestIntegrationKillRefused(t *testing.T) { + t.Parallel() + + env := newTestEnvWithOper(t) + + alice := env.dial(t) + alice.register("alice") + + bob := env.dial(t) + bob.register("bob") + + for _, step := range []struct { + line, numeric string + }{ + {"KILL bob :nope", " 481 "}, + {"OPER testoper testpass", " 381 "}, + {"KILL nobody123 :gone", " 401 "}, + {"KILL alice :me", " 483 "}, + } { + alice.send(step.line) + + reply := alice.readUntil(func(l string) bool { + return strings.Contains(l, step.numeric) + }) + assertContains(t, reply, step.numeric, step.line) + } +} + +// TestIntegrationKill checks that the victim of a KILL is +// told why, is disconnected, and is gone from its channels. +func TestIntegrationKill(t *testing.T) { + t.Parallel() + + env := newTestEnvWithOper(t) + + alice := env.dial(t) + alice.register("alice") + + bob := env.dial(t) + bob.register("bob") + + alice.joinAndDrain("#killtest") + bob.joinAndDrain("#killtest") + + // Drain alice's view of bob's join. + alice.readUntil(func(l string) bool { + return strings.Contains(l, "JOIN") && + strings.Contains(l, "bob") + }) + + alice.becomeOper() + alice.send("KILL bob :bad behavior") + + bobLines := bob.readUntilClosed() + assertContains( + t, bobLines, "KILL", + "victim receives KILL before disconnect", + ) + assertContains( + t, bobLines, "ERROR :Closing Link", + "victim receives ERROR before disconnect", + ) + assertContains( + t, bobLines, "bad behavior", + "KILL reason delivered to victim", + ) + + // alice should see bob's QUIT relay. + aliceSeesQuit := alice.readUntil(func(l string) bool { + return strings.Contains(l, "QUIT") && + strings.Contains(l, "bob") + }) + assertContains( + t, aliceSeesQuit, "Killed", + "KILL reason in QUIT message", + ) + + // bob must be gone from the channel member list. + alice.send("NAMES #killtest") + + aliceNames := alice.readUntil(func(l string) bool { + return strings.Contains(l, " 366 ") + }) + assertContains( + t, aliceNames, "alice", + "alice still in NAMES after killing bob", + ) + assertNotContains( + t, aliceNames, "bob", + "killed user must not appear in NAMES", + ) + + // ...nor from WHO. + alice.send("WHO #killtest") + + aliceWho := alice.readUntil(func(l string) bool { + return strings.Contains(l, " 315 ") + }) + assertNotContains( + t, aliceWho, "bob", + "killed user must not appear in WHO", + ) +} + +// TestIntegrationWallops verifies the WALLOPS command: +// oper can broadcast to +w users. +func TestIntegrationWallops(t *testing.T) { + t.Parallel() + + env := newTestEnvWithOper(t) + + alice := env.dial(t) + alice.register("alice") + + bob := env.dial(t) + bob.register("bob") + + // Non-oper WALLOPS should fail. + alice.send("WALLOPS :test broadcast") + + aliceWallopsFail := alice.readUntil(func(l string) bool { + return strings.Contains(l, " 481 ") + }) + assertContains( + t, aliceWallopsFail, " 481 ", + "ERR_NOPRIVILEGES for non-oper WALLOPS", + ) + + alice.becomeOper() + + // bob sets +w to receive wallops. + bob.send("MODE bob +w") + bob.readUntil(func(l string) bool { + return strings.Contains(l, " 221 ") + }) + + // alice sends WALLOPS. + alice.send("WALLOPS :important announcement") + + // bob (who has +w) should receive it. + bobWallops := bob.readUntil(func(l string) bool { + return strings.Contains( + l, "important announcement", + ) + }) + assertContains( + t, bobWallops, "important announcement", + "bob receives WALLOPS message", + ) + assertContains( + t, bobWallops, "WALLOPS", + "message is WALLOPS command", + ) +} + +// TestIntegrationUserMode covers MODE for a nick: changes +// and queries of your own modes in any letter case, refusal +// for another nick, and a rejected mode string changing +// nothing. +func TestIntegrationUserMode(t *testing.T) { + t.Parallel() + + env := newTestEnv(t) + + alice := env.dial(t) + alice.register("alice") + + bob := env.dial(t) + bob.register("bob") + + const ( + notYours = " 502 alice :Can't change mode for other users" + unknown = " 501 alice :Unknown MODE flag" + ) + + for _, step := range []struct { + line, want string + }{ + {"MODE ALICE +w", " 221 alice +w"}, + {"MODE bob", notYours}, + {"MODE bob -w", notYours}, + {"MODE alice xw", unknown}, + {"MODE alice -w+z", unknown}, + {"MODE alice", " 221 alice +w"}, + {"MODE alice +w-w", " 221 alice +"}, + } { + alice.send(step.line) + + reply := alice.readUntil(func(l string) bool { + return strings.Contains(l, " 221 ") || + strings.Contains(l, " 501 ") || + strings.Contains(l, " 502 ") + }) + + last := reply[len(reply)-1] + if !strings.HasSuffix(last, step.want) { + t.Errorf("%s: want %q, got %q", step.line, step.want, last) + } + } +} + // TestIntegrationModeSecret tests +s (secret) channel // mode — verifies that +s can be set and the mode is // reflected in MODE queries. @@ -915,3 +1282,30 @@ func TestIntegrationThirdClientObserver(t *testing.T) { "carol receives trio message", ) } + +// TestIntegrationDefaultServerNameFallback checks that with +// SERVER_NAME unset, as it is by default, VERSION, ADMIN and +// TIME name the server "neoirc". +func TestIntegrationDefaultServerNameFallback(t *testing.T) { + t.Parallel() + + env := newTestEnvWithConfig(t, &config.Config{}) + + alice := env.dial(t) + alice.register("alice") + + for _, step := range []struct { + line, lastNumeric, want string + }{ + {"VERSION", " 351 ", " 351 alice neoirc-test. neoirc "}, + {"ADMIN", " 259 ", " 256 alice neoirc "}, + {"TIME", " 391 ", " 391 alice neoirc "}, + } { + alice.send(step.line) + + reply := alice.readUntil(func(l string) bool { + return strings.Contains(l, step.lastNumeric) + }) + assertContains(t, reply, step.want, step.line) + } +} diff --git a/internal/ircserver/relay.go b/internal/ircserver/relay.go index 6d4195f..60ccb74 100644 --- a/internal/ircserver/relay.go +++ b/internal/ircserver/relay.go @@ -120,6 +120,8 @@ func (c *Conn) deliverIRCMessage( c.deliverKickMsg(msg, text) case command == "INVITE": c.deliverInviteMsg(msg, text) + case command == irc.CmdWallops: + c.deliverWallops(msg, text) case command == irc.CmdMode: c.deliverMode(msg, text) case command == irc.CmdPing: @@ -337,6 +339,18 @@ func (c *Conn) deliverInviteMsg( c.sendFromServer("NOTICE", nick, text) } +// deliverWallops sends a WALLOPS notification. +func (c *Conn) deliverWallops( + msg *db.IRCMessage, + text string, +) { + prefix := msg.From + "!" + msg.From + "@*" + + c.send(FormatMessage( + prefix, irc.CmdWallops, text, + )) +} + // deliverMode sends a MODE change notification. func (c *Conn) deliverMode( msg *db.IRCMessage, diff --git a/internal/ircserver/server_test.go b/internal/ircserver/server_test.go index 0ea6f4e..3717474 100644 --- a/internal/ircserver/server_test.go +++ b/internal/ircserver/server_test.go @@ -36,9 +36,39 @@ type testEnv struct { srv *ircserver.Server } +// testServerName is the SERVER_NAME of newTestEnv. +const testServerName = "test.irc" + func newTestEnv(t *testing.T) *testEnv { t.Helper() + return newTestEnvWithConfig(t, &config.Config{ + ServerName: testServerName, + MOTD: "Welcome to test IRC", + }) +} + +// newTestEnvWithOper is newTestEnv with the operator name +// testoper and password testpass. +func newTestEnvWithOper(t *testing.T) *testEnv { + t.Helper() + + return newTestEnvWithConfig(t, &config.Config{ + ServerName: testServerName, + MOTD: "Welcome to test IRC", + OperName: "testoper", + OperPassword: "testpass", + }) +} + +// newTestEnvWithConfig starts an IRC server with cfg on a +// fresh database. +func newTestEnvWithConfig( + t *testing.T, + cfg *config.Config, +) *testEnv { + t.Helper() + dsn := fmt.Sprintf( "file:%s?mode=memory&cache=shared&_journal_mode=WAL", t.Name(), @@ -67,11 +97,6 @@ func newTestEnv(t *testing.T) *testEnv { brk := broker.New() - cfg := &config.Config{ //nolint:exhaustruct - ServerName: "test.irc", - MOTD: "Welcome to test IRC", - } - var listenConfig net.ListenConfig listener, err := listenConfig.Listen(t.Context(), "tcp", "127.0.0.1:0") @@ -215,6 +240,33 @@ func (tc *testClient) register(nick string) []string { }) } +// readUntilClosed returns the lines received until the +// server closes the connection, and fails the test if it is +// still open after testTimeout. +func (tc *testClient) readUntilClosed() []string { + tc.t.Helper() + + _ = tc.conn.SetReadDeadline( + time.Now().Add(testTimeout), + ) + + var lines []string + + for tc.scanner.Scan() { + lines = append(lines, tc.scanner.Text()) + } + + err := tc.scanner.Err() + if err != nil { + tc.t.Fatalf( + "connection not closed: %v (lines: %v)", + err, lines, + ) + } + + return lines +} + // assertContains checks that at least one line matches the // given substring. func assertContains( @@ -233,6 +285,27 @@ func assertContains( t.Errorf("did not find %q in output: %s", substr, description) } +// assertNotContains checks that no line matches the given +// substring. +func assertNotContains( + t *testing.T, + lines []string, + substr, description string, +) { + t.Helper() + + for _, line := range lines { + if strings.Contains(line, substr) { + t.Errorf( + "unexpectedly found %q in output: %s", + substr, description, + ) + + return + } + } +} + // joinAndDrain joins a channel and reads until // RPL_ENDOFNAMES. func (tc *testClient) joinAndDrain(channel string) { diff --git a/internal/service/service.go b/internal/service/service.go index 690ce08..16e2fa4 100644 --- a/internal/service/service.go +++ b/internal/service/service.go @@ -5,16 +5,21 @@ package service import ( "context" "crypto/subtle" + "database/sql" "encoding/json" + "errors" "fmt" "log/slog" "strconv" "strings" + "sync" + "time" "go.uber.org/fx" "sneak.berlin/go/neoirc/internal/broker" "sneak.berlin/go/neoirc/internal/config" "sneak.berlin/go/neoirc/internal/db" + "sneak.berlin/go/neoirc/internal/globals" "sneak.berlin/go/neoirc/internal/logger" "sneak.berlin/go/neoirc/pkg/irc" ) @@ -22,9 +27,14 @@ import ( // Error texts that several commands reply with. const ( msgNoSuchChannel = "No such channel" + msgNoSuchNick = "No such nick/channel" msgNotChannelOp = "You're not channel operator" ) +// maxUserhostNicks is how many nicks one USERHOST answers +// for (RFC 2812). +const maxUserhostNicks = 5 + // Params defines the dependencies for creating a Service. type Params struct { fx.In @@ -33,23 +43,38 @@ type Params struct { Config *config.Config Database *db.Database Broker *broker.Broker + Globals *globals.Globals +} + +// WireConn is a registered IRC connection, which KILL closes. +// HTTP clients hold no connection and register none. +type WireConn interface { + // Disconnect tells the client why and closes the + // connection. + Disconnect(reason string) } // Service provides shared business logic for IRC commands. type Service struct { - db *db.Database - broker *broker.Broker - config *config.Config - log *slog.Logger + db *db.Database + broker *broker.Broker + config *config.Config + globals *globals.Globals + log *slog.Logger + + wireMu sync.Mutex + wireConns map[int64]WireConn } // New creates a new Service. func New(params Params) *Service { return &Service{ - db: params.Database, - broker: params.Broker, - config: params.Config, - log: params.Logger.Get(), + db: params.Database, + broker: params.Broker, + config: params.Config, + globals: params.Globals, + log: params.Logger.Get(), + wireConns: make(map[int64]WireConn), } } @@ -59,16 +84,205 @@ func NewTestService( database *db.Database, brk *broker.Broker, cfg *config.Config, + globs *globals.Globals, log *slog.Logger, ) *Service { return &Service{ - db: database, - broker: brk, - config: cfg, - log: log, + db: database, + broker: brk, + config: cfg, + globals: globs, + log: log, + wireConns: make(map[int64]WireConn), } } +// ServerVersion returns the version that VERSION and INFO +// report on both transports, such as "neoirc-1.2.3". +func (s *Service) ServerVersion() string { + version := s.globals.Version + if version == "" { + version = "dev" + } + + return s.globals.Appname + "-" + version +} + +// InfoLines returns the lines that INFO sends on both +// transports. +func (s *Service) InfoLines() []string { + return []string{ + "neoirc — IRC semantics over HTTP", + "Version: " + s.ServerVersion(), + "Written in Go", + "Started: " + s.globals.StartTime.Format(time.RFC1123), + } +} + +// UserhostReply returns the RPL_USERHOST text for the first +// five of the given nicks: "nick=+user@host" entries joined +// by spaces, with * after the nick of an operator and - in +// place of + for a user who is away. Nicks with no session +// are left out. An empty hostname is reported as +// serverName. +func (s *Service) UserhostReply( + ctx context.Context, + nicks []string, + serverName string, +) (string, error) { + if len(nicks) > maxUserhostNicks { + nicks = nicks[:maxUserhostNicks] + } + + infos, err := s.db.GetUserhostInfo(ctx, nicks) + if err != nil { + return "", fmt.Errorf("userhost: %w", err) + } + + replies := make([]string, 0, len(infos)) + + for _, info := range infos { + operStar := "" + if info.IsOper { + operStar = "*" + } + + away := "+" + if info.AwayMessage != "" { + away = "-" + } + + username := info.Username + if username == "" { + username = info.Nick + } + + hostname := info.Hostname + if hostname == "" { + hostname = serverName + } + + replies = append(replies, + info.Nick+operStar+"="+away+username+"@"+hostname, + ) + } + + return strings.Join(replies, " "), nil +} + +// RegisterWireConn records the IRC connection of a session, +// so that KillUser can close it. +func (s *Service) RegisterWireConn( + sessionID int64, + conn WireConn, +) { + s.wireMu.Lock() + defer s.wireMu.Unlock() + + s.wireConns[sessionID] = conn +} + +// UnregisterWireConn removes what RegisterWireConn recorded, +// unless the session has since registered another +// connection. +func (s *Service) UnregisterWireConn( + sessionID int64, + conn WireConn, +) { + s.wireMu.Lock() + defer s.wireMu.Unlock() + + if s.wireConns[sessionID] == conn { + delete(s.wireConns, sessionID) + } +} + +// KillUser carries out an operator's KILL: the target's +// channel peers see it quit, its session is deleted, and +// its IRC connection, if it has one, is closed. +func (s *Service) KillUser( + ctx context.Context, + sessionID int64, + nick, targetNick, reason string, +) error { + err := s.requireOper(ctx, sessionID) + if err != nil { + return err + } + + targetSID, err := s.db.GetSessionByNick(ctx, targetNick) + if errors.Is(err, sql.ErrNoRows) { + return &IRCError{ + irc.ErrNoSuchNick, + []string{targetNick}, + msgNoSuchNick, + } + } + + if err != nil { + return fmt.Errorf("kill: %w", err) + } + + if targetSID == sessionID { + return &IRCError{ + irc.ErrCantKillServer, + nil, + "You cannot KILL yourself", + } + } + + quitReason := "Killed (" + nick + " (" + reason + "))" + + s.BroadcastQuit(ctx, targetSID, targetNick, quitReason) + + s.wireMu.Lock() + conn := s.wireConns[targetSID] + s.wireMu.Unlock() + + if conn != nil { + conn.Disconnect(quitReason) + } + + return nil +} + +// SendWallops carries out an operator's WALLOPS: the +// message goes to every user with user mode +w. +func (s *Service) SendWallops( + ctx context.Context, + sessionID int64, + nick, message string, +) error { + err := s.requireOper(ctx, sessionID) + if err != nil { + return err + } + + recipients, err := s.db.GetWallopsSessionIDs(ctx) + if err != nil { + return fmt.Errorf("wallops: %w", err) + } + + if len(recipients) == 0 { + return nil + } + + body, err := json.Marshal([]string{message}) + if err != nil { + return fmt.Errorf("wallops: %w", err) + } + + _, _, err = s.FanOut( + ctx, irc.CmdWallops, nick, "*", + nil, body, nil, recipients, + ) + if err != nil { + return fmt.Errorf("wallops: %w", err) + } + + return nil +} + // IRCError represents an IRC protocol-level error with a // numeric code that both transports can map to responses. type IRCError struct { @@ -437,7 +651,7 @@ func (s *Service) KickUser( return &IRCError{ irc.ErrNoSuchNick, []string{targetNick}, - "No such nick/channel", + msgNoSuchNick, } } @@ -649,7 +863,7 @@ func (s *Service) ApplyMemberMode( return &IRCError{ irc.ErrNoSuchNick, []string{targetNick}, - "No such nick/channel", + msgNoSuchNick, } } @@ -792,6 +1006,139 @@ func (s *Service) QueryChannelMode( return modes + modeParams } +// QueryUserMode returns the session's user mode string, +// such as "+", "+w" or "+ow". +func (s *Service) QueryUserMode( + ctx context.Context, + sessionID int64, +) (string, error) { + modes := "+" + + isOper, err := s.db.IsSessionOper(ctx, sessionID) + if err != nil { + return "", fmt.Errorf( + "query oper flag: %w", err, + ) + } + + if isOper { + modes += "o" + } + + isWallops, err := s.db.IsSessionWallops( + ctx, sessionID, + ) + if err != nil { + return "", fmt.Errorf( + "query wallops flag: %w", err, + ) + } + + if isWallops { + modes += "w" + } + + return modes, nil +} + +// ApplyUserMode applies a user mode string to the session +// and returns the resulting mode string. The whole string +// is parsed before anything is written, and the flags are +// written in one transaction, so a change that is rejected +// or fails leaves the modes as they were. +func (s *Service) ApplyUserMode( + ctx context.Context, + sessionID int64, + modeStr string, +) (string, error) { + wallops, oper, err := parseUserModeString(modeStr) + if err != nil { + return "", err + } + + err = s.db.SetSessionUserModes( + ctx, sessionID, wallops, oper, + ) + if err != nil { + return "", fmt.Errorf("apply user modes: %w", err) + } + + return s.QueryUserMode(ctx, sessionID) +} + +// parseUserModeString parses a user mode string such as +// "+w", "-o" or "+w-o" into the new value of each flag it +// names, or nil for a flag it does not name; a later letter +// overrides an earlier one. The string must start with + or +// - and name at least one mode. w may be set or unset, o +// only unset (OPER sets it). Anything else is rejected with +// ERR_UMODEUNKNOWNFLAG. +func parseUserModeString( + modeStr string, +) (*bool, *bool, error) { + unknownFlag := &IRCError{ + irc.ErrUmodeUnknownFlag, nil, "Unknown MODE flag", + } + + if !strings.HasPrefix(modeStr, "+") && + !strings.HasPrefix(modeStr, "-") { + return nil, nil, unknownFlag + } + + var wallops, oper *bool + + adding := true + + for _, modeChar := range modeStr { + switch modeChar { + case '+': + adding = true + case '-': + adding = false + case 'w': + value := adding + wallops = &value + case 'o': + if adding { + return nil, nil, unknownFlag + } + + value := false + oper = &value + default: + return nil, nil, unknownFlag + } + } + + if wallops == nil && oper == nil { + return nil, nil, unknownFlag + } + + return wallops, oper, nil +} + +// requireOper returns ERR_NOPRIVILEGES unless the session +// is a server operator. +func (s *Service) requireOper( + ctx context.Context, + sessionID int64, +) error { + isOper, err := s.db.IsSessionOper(ctx, sessionID) + if err != nil { + return fmt.Errorf("check oper: %w", err) + } + + if !isOper { + return &IRCError{ + irc.ErrNoPrivileges, + nil, + "Permission Denied- You're not an IRC operator", + } + } + + return nil +} + // broadcastNickChange notifies channel peers of a nick // change. func (s *Service) broadcastNickChange( diff --git a/internal/service/service_test.go b/internal/service/service_test.go index 864a8a8..933235d 100644 --- a/internal/service/service_test.go +++ b/internal/service/service_test.go @@ -10,7 +10,10 @@ import ( "errors" "fmt" "os" + "slices" + "strings" "testing" + "time" "go.uber.org/fx" "go.uber.org/fx/fxtest" @@ -55,9 +58,10 @@ func newTestEnv(t *testing.T) *testEnv { app := fxtest.New(t, fx.Provide( func() *globals.Globals { - return &globals.Globals{ //nolint:exhaustruct - Appname: "neoirc-test", - Version: "test", + return &globals.Globals{ + Appname: "neoirc-test", + Version: "test", + StartTime: time.Now(), } }, logger.New, @@ -363,3 +367,229 @@ func TestSendChannelMessage_Moderated(t *testing.T) { t.Errorf("operator should be able to send in moderated channel: %v", err) } } + +func TestQueryUserMode(t *testing.T) { + env := newTestEnv(t) + ctx := t.Context() + + sid := createSession(ctx, t, env.db, "alice") + + modes, err := env.svc.QueryUserMode(ctx, sid) + if err != nil { + t.Fatalf("query user mode: %v", err) + } + + if modes != "+" { + t.Errorf("expected +, got %s", modes) + } + + setUserModes(ctx, t, env.db, sid, userModes{wallops: true}) + + modes, err = env.svc.QueryUserMode(ctx, sid) + if err != nil { + t.Fatalf("query user mode: %v", err) + } + + if modes != "+w" { + t.Errorf("expected +w, got %s", modes) + } + + setUserModes( + ctx, t, env.db, sid, userModes{oper: true, wallops: true}, + ) + + modes, err = env.svc.QueryUserMode(ctx, sid) + if err != nil { + t.Fatalf("query user mode: %v", err) + } + + if modes != "+ow" { + t.Errorf("expected +ow, got %s", modes) + } +} + +func TestUserhostReply(t *testing.T) { + env := newTestEnv(t) + ctx := t.Context() + + createSession(ctx, t, env.db, "alice") + bobID := createSession(ctx, t, env.db, "bob") + + // No username or hostname, and an operator. + operID, _, _, err := env.db.CreateSession( + ctx, "oper", "", "", "", + ) + if err != nil { + t.Fatal(err) + } + + setUserModes(ctx, t, env.db, operID, userModes{oper: true}) + + _, err = env.svc.SetAway(ctx, bobID, "gone fishing") + if err != nil { + t.Fatal(err) + } + + reply, err := env.svc.UserhostReply(ctx, []string{ + "alice", "bob", "nobody", "oper", + }, "srv") + if err != nil { + t.Fatal(err) + } + + want := "alice=+alice@localhost bob=-bob@localhost " + + "oper*=+oper@srv" + if reply != want { + t.Errorf("want %q, got %q", want, reply) + } + + reply, err = env.svc.UserhostReply( + ctx, slices.Repeat([]string{"bob"}, 6), "srv", + ) + if err != nil { + t.Fatal(err) + } + + if strings.Count(reply, "bob=") != 5 { + t.Errorf("want 5 entries for 6 nicks, got %q", reply) + } +} + +// userModes is the stored state of a session's user mode +// flags. +type userModes struct { + oper bool + wallops bool +} + +// applyUserModeTest is one mode string, the flags stored +// before it is applied, and the expected outcome. An empty +// wantModes means the string is rejected with +// ERR_UMODEUNKNOWNFLAG, and then after must equal before. +type applyUserModeTest struct { + modeStr string + before userModes + wantModes string + after userModes +} + +func applyUserModeTests() []applyUserModeTest { + none := userModes{} + oper := userModes{oper: true} + wallops := userModes{wallops: true} + both := userModes{oper: true, wallops: true} + + return []applyUserModeTest{ + {"+w", none, "+w", wallops}, + {"-w", wallops, "+", none}, + {"-o", oper, "+", none}, + {"-wo", both, "+", none}, + {"+w-o", oper, "+w", wallops}, + {"-w+w", none, "+w", wallops}, + {"+w-w", wallops, "+", none}, + {"-w+o", both, "", both}, + {"+o-w+w", both, "", both}, + {"+wo", none, "", none}, + {"+wz", none, "", none}, + {"+z", none, "", none}, + {"-x+y", none, "", none}, + {"+y-x", none, "", none}, + {"w", none, "", none}, + {"xw", wallops, "", wallops}, + {"", none, "", none}, + {"+", none, "", none}, + {"-", none, "", none}, + {"+-+", none, "", none}, + } +} + +func TestApplyUserMode(t *testing.T) { + env := newTestEnv(t) + + for i, test := range applyUserModeTests() { + t.Run(test.modeStr, func(t *testing.T) { + ctx := t.Context() + sid := createSession( + ctx, t, env.db, fmt.Sprintf("user%d", i), + ) + setUserModes(ctx, t, env.db, sid, test.before) + + modes, err := env.svc.ApplyUserMode( + ctx, sid, test.modeStr, + ) + checkApplyUserModeResult(t, test.wantModes, modes, err) + + got := getUserModes(ctx, t, env.db, sid) + if got != test.after { + t.Errorf( + "stored modes: want %+v, got %+v", + test.after, got, + ) + } + }) + } +} + +func checkApplyUserModeResult( + t *testing.T, + wantModes, modes string, + err error, +) { + t.Helper() + + if wantModes == "" { + var ircErr *service.IRCError + if !errors.As(err, &ircErr) || + ircErr.Code != irc.ErrUmodeUnknownFlag { + t.Fatalf("want ERR_UMODEUNKNOWNFLAG, got %v", err) + } + + return + } + + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + + if modes != wantModes { + t.Errorf("modes: want %q, got %q", wantModes, modes) + } +} + +func setUserModes( + ctx context.Context, + t *testing.T, + database *db.Database, + sid int64, + modes userModes, +) { + t.Helper() + + err := database.SetSessionUserModes( + ctx, sid, &modes.wallops, &modes.oper, + ) + if err != nil { + t.Fatalf("set user modes: %v", err) + } +} + +func getUserModes( + ctx context.Context, + t *testing.T, + database *db.Database, + sid int64, +) userModes { + t.Helper() + + oper, err := database.IsSessionOper(ctx, sid) + if err != nil { + t.Fatalf("read oper: %v", err) + } + + wallops, err := database.IsSessionWallops(ctx, sid) + if err != nil { + t.Fatalf("read wallops: %v", err) + } + + return userModes{oper: oper, wallops: wallops} +} diff --git a/pkg/irc/commands.go b/pkg/irc/commands.go index 60ef85e..4e2e8af 100644 --- a/pkg/irc/commands.go +++ b/pkg/irc/commands.go @@ -2,26 +2,33 @@ package irc // IRC command names (RFC 1459 / RFC 2812). const ( - CmdAway = "AWAY" - CmdInvite = "INVITE" - CmdJoin = "JOIN" - CmdKick = "KICK" - CmdList = "LIST" - CmdLusers = "LUSERS" - CmdMode = "MODE" - CmdMotd = "MOTD" - CmdNames = "NAMES" - CmdNick = "NICK" - CmdNotice = "NOTICE" - CmdOper = "OPER" - CmdPass = "PASS" - CmdPart = "PART" - CmdPing = "PING" - CmdPong = "PONG" - CmdPrivmsg = "PRIVMSG" - CmdQuit = "QUIT" - CmdTopic = "TOPIC" - CmdUser = "USER" - CmdWho = "WHO" - CmdWhois = "WHOIS" + CmdAdmin = "ADMIN" + CmdAway = "AWAY" + CmdInfo = "INFO" + CmdInvite = "INVITE" + CmdJoin = "JOIN" + CmdKick = "KICK" + CmdKill = "KILL" + CmdList = "LIST" + CmdLusers = "LUSERS" + CmdMode = "MODE" + CmdMotd = "MOTD" + CmdNames = "NAMES" + CmdNick = "NICK" + CmdNotice = "NOTICE" + CmdOper = "OPER" + CmdPass = "PASS" + CmdPart = "PART" + CmdPing = "PING" + CmdPong = "PONG" + CmdPrivmsg = "PRIVMSG" + CmdQuit = "QUIT" + CmdTime = "TIME" + CmdTopic = "TOPIC" + CmdUser = "USER" + CmdUserhost = "USERHOST" + CmdVersion = "VERSION" + CmdWallops = "WALLOPS" + CmdWho = "WHO" + CmdWhois = "WHOIS" )