diff --git a/README.md b/README.md index 83ee55d..44f56c0 100644 --- a/README.md +++ b/README.md @@ -9,7 +9,8 @@ people use (CIDER, Calva, Conjure, vim-fireplace, REPLy and so on). It talks to the server over a socket like any client would, runs a set of checks against it and tells you what's broken and which clients it breaks. It can also check the other side of the conversation, i.e. the requests -an nREPL client sends. +an nREPL client sends, and stand in for other servers in a client's +tests. The [nREPL protocol spec](https://spec.nrepl.org) is still a draft and in a few places it disagrees with what clients actually do. When that @@ -89,6 +90,21 @@ and that sessions get closed in the end. Here a failure means that some server won't work properly with your client, and the report links to the server code in question. +To see how your client deals with the replies of different servers, run +its tests against `proof serve`. It's a small nREPL server that behaves +like nREPL itself, unless you ask it to behave like some other server in +one particular way: + +```shell +$ proof serve -listen 127.0.0.1:7888 last-value no-err +``` + +Here it sends only the value of the last form (like Basilisp, jank and +dialtone) and drops what the code prints to stderr (like Basilisp, +dialtone and repartee). `proof list` shows all the scenarios, and every +one of them is something a real server does (or something TCP can do to +the replies). + ## Documentation - [Usage](doc/usage.md) - checking your server, reading the report, @@ -106,8 +122,8 @@ server code in question. proof is still in its early days. Right now it covers the core of the protocol (`describe`, unknown ops, `eval`, sessions, `stdin` and the wire format), -along with the requests clients send, and `proof list` will show you all -the checks. +along with the requests clients send and the server differences clients +have to deal with, and `proof list` will show you all the checks. Here's what's coming next: @@ -116,9 +132,6 @@ Here's what's coming next: clients disconnecting in the middle of an evaluation) - replaying what real clients send (e.g. when CIDER or Calva connect to a server) as client profiles -- a server that misbehaves on purpose (late output, output split into - many messages and so on), so client test suites can check how their - client deals with replies - publishing the compatibility matrix somewhere nicer than a CI job summary diff --git a/cmd/proof/main.go b/cmd/proof/main.go index f1c2594..b100a86 100644 --- a/cmd/proof/main.go +++ b/cmd/proof/main.go @@ -1,5 +1,5 @@ // Command proof checks an nREPL server's compatibility with the clients -// people actually use, and what a client sends to a server. +// people actually use, and checks clients against servers. package main import ( @@ -17,6 +17,7 @@ import ( "github.com/nrepl/proof/internal/checks" "github.com/nrepl/proof/internal/profile" "github.com/nrepl/proof/internal/report" + "github.com/nrepl/proof/internal/serve" "github.com/nrepl/proof/internal/server" ) @@ -25,17 +26,18 @@ const version = "0.1.0-dev" const usage = `proof checks an nREPL server's compatibility with existing clients. Usage: - proof run [flags] PROFILE run the checks against the server a profile describes - proof proxy [flags] [PROFILE] check what a client sends to a server, by sitting between them - proof matrix REPORT... build a Markdown compatibility matrix from JSON reports - proof list list every check and rule + proof run [flags] PROFILE run the checks against the server a profile describes + proof proxy [flags] [PROFILE] check what a client sends to a server, by sitting between them + proof serve [flags] [SCENARIO...] be a server for a client's tests, acting like other servers where asked + proof matrix REPORT... build a Markdown compatibility matrix from JSON reports + proof list list every check, rule and scenario proof version The exit status of run is 0 when everything passed (or failed as the profile expects), 1 when the server failed checks, 2 when proof couldn't start (bad flags or profile, or the server didn't come up), and 3 when some -checks couldn't run at all. The same goes for proxy, where 1 means the -client failed rules and 3 means no client sent anything. +checks couldn't run at all. The same goes for proxy and serve, where 1 +means the client failed rules and 3 means no client sent anything. Run flags: ` @@ -50,6 +52,8 @@ func main() { os.Exit(run(os.Args[2:])) case "proxy": os.Exit(runProxy(os.Args[2:])) + case "serve": + os.Exit(runServe(os.Args[2:])) case "matrix": os.Exit(matrix(os.Args[2:])) case "list": @@ -69,6 +73,8 @@ func printUsage(w io.Writer) { runFlags(w).PrintDefaults() fmt.Fprint(w, "\nProxy flags:\n") proxyFlags(w, &proxyOptions{}).PrintDefaults() + fmt.Fprint(w, "\nServe flags:\n") + serveFlags(w, &clientOptions{}).PrintDefaults() } type options struct { @@ -278,4 +284,8 @@ func list(w io.Writer) { for _, e := range append(catalog(), ruleEntries(checks.ClientRules())...) { fmt.Fprintf(w, "%-28s %-4s %s\n", e.id, e.severity, e.title) } + fmt.Fprint(w, "\nScenarios for proof serve:\n") + for _, s := range serve.Scenarios() { + fmt.Fprintf(w, "%-28s %s (%s)\n", s.Name, s.Title, s.Who) + } } diff --git a/cmd/proof/proxy.go b/cmd/proof/proxy.go index 1198469..492eb16 100644 --- a/cmd/proof/proxy.go +++ b/cmd/proof/proxy.go @@ -19,23 +19,39 @@ import ( "github.com/nrepl/proof/internal/server" ) -type proxyOptions struct { - address string +// clientOptions are the options of the commands that check clients. +type clientOptions struct { listen string json string verbose bool } +func (o *clientOptions) register(fs *flag.FlagSet) { + fs.StringVar(&o.listen, "listen", "127.0.0.1:0", "accept clients on `host:port` (port 0 picks a free one)") + fs.StringVar(&o.json, "json", "", "also write a JSON report to `file`") + fs.BoolVar(&o.verbose, "v", false, "show every message the client and the server exchanged") +} + +type proxyOptions struct { + clientOptions + address string +} + func proxyFlags(out io.Writer, o *proxyOptions) *flag.FlagSet { fs := flag.NewFlagSet("proxy", flag.ContinueOnError) fs.SetOutput(out) fs.StringVar(&o.address, "address", "", "forward clients to a server already running at `host:port` instead of launching the one the profile describes") - fs.StringVar(&o.listen, "listen", "127.0.0.1:0", "accept clients on `host:port` (port 0 picks a free one)") - fs.StringVar(&o.json, "json", "", "also write a JSON report to `file`") - fs.BoolVar(&o.verbose, "v", false, "show every message the client and the server exchanged") + o.register(fs) return fs } +// logTo writes log messages to w, one per line. +func logTo(w io.Writer) func(format string, args ...any) { + return func(format string, args ...any) { + fmt.Fprintf(w, format+"\n", args...) + } +} + // runProxy sits between a client and a server until it's interrupted, and // then grades everything the client sent. func runProxy(args []string) int { @@ -95,9 +111,7 @@ func proxyUntil(ctx context.Context, args []string, stdout, stderr io.Writer, li fmt.Fprintln(stderr, "proof:", err) return 2 } - px.Logf = func(format string, args ...any) { - fmt.Fprintf(stderr, format+"\n", args...) - } + px.Logf = logTo(stderr) go px.Serve() started := time.Now() fmt.Fprintf(stderr, "Forwarding %s to %s. Connect your client to %s and press Ctrl-C when it's done.\n", @@ -108,19 +122,19 @@ func proxyUntil(ctx context.Context, args []string, stdout, stderr io.Writer, li <-ctx.Done() fmt.Fprintln(stderr) - traffic := px.Stop(time.Second) + r := report.Run{Proof: version, Server: "client traffic to " + name, Address: upstream, Started: started} + return gradeClients(stdout, stderr, r, px.Stop(time.Second), o.clientOptions, srv) +} + +// gradeClients grades the requests in traffic, reports on them, and +// returns the exit status. srv is the server proof started, if any. +func gradeClients(stdout, stderr io.Writer, r report.Run, traffic []check.Traffic, o clientOptions, srv *server.Server) int { if len(traffic) == 0 { fmt.Fprintln(stderr, "proof: no client sent anything to the server, so there's nothing to check") showDeath(stderr, srv) return 3 } - r := report.Run{ - Proof: version, - Server: "client traffic to " + name, - Address: upstream, - Started: started, - Results: check.Grade(checks.ClientRules(), traffic), - } + r.Results = check.Grade(checks.ClientRules(), traffic) report.Text(stdout, r, o.verbose) if o.verbose { for _, tr := range traffic { diff --git a/cmd/proof/proxy_test.go b/cmd/proof/proxy_test.go index 94b4e00..987a4de 100644 --- a/cmd/proof/proxy_test.go +++ b/cmd/proof/proxy_test.go @@ -50,9 +50,13 @@ func doneServer(t *testing.T) string { return ln.Addr().String() } -// proxyRun runs proof proxy with a client that sends the given frames and -// hangs up, and returns the exit status and output. -func proxyRun(t *testing.T, frames []string, args ...string) (int, string, string) { +// command is how the tests run proof proxy and proof serve: until ctx is +// done, telling listening where clients go. +type command func(ctx context.Context, args []string, stdout, stderr io.Writer, listening func(addr string)) int + +// runWithClient runs a command with a client that sends the given frames +// and hangs up, and returns the exit status and output. +func runWithClient(t *testing.T, run command, frames []string, args ...string) (int, string, string) { t.Helper() ctx, cancel := context.WithCancel(context.Background()) defer cancel() @@ -75,11 +79,17 @@ func proxyRun(t *testing.T, frames []string, args ...string) (int, string, strin io.Copy(io.Discard, c) } var stdout, stderr bytes.Buffer - args = append([]string{"-address", doneServer(t)}, args...) - code := proxyUntil(ctx, args, &stdout, &stderr, client) + code := run(ctx, args, &stdout, &stderr, client) return code, stdout.String(), stderr.String() } +// proxyRun runs proof proxy in front of a server that answers everything +// with done. +func proxyRun(t *testing.T, frames []string, args ...string) (int, string, string) { + t.Helper() + return runWithClient(t, proxyUntil, frames, append([]string{"-address", doneServer(t)}, args...)...) +} + func TestProxyExitCodes(t *testing.T) { cases := []struct { name string @@ -129,10 +139,12 @@ func TestProxyNeedsAServer(t *testing.T) { } } -func TestListShowsClientRules(t *testing.T) { +func TestListShowsEverything(t *testing.T) { var buf bytes.Buffer list(&buf) - if !strings.Contains(buf.String(), "client.need-input") || !strings.Contains(buf.String(), "wire.dict") { - t.Errorf("list is missing rules:\n%s", buf.String()) + for _, want := range []string{"eval.value", "wire.dict", "client.need-input", "split-output"} { + if !strings.Contains(buf.String(), want) { + t.Errorf("list is missing %s:\n%s", want, buf.String()) + } } } diff --git a/cmd/proof/serve.go b/cmd/proof/serve.go new file mode 100644 index 0000000..f4bbd54 --- /dev/null +++ b/cmd/proof/serve.go @@ -0,0 +1,64 @@ +package main + +import ( + "context" + "flag" + "fmt" + "io" + "os" + "os/signal" + "strings" + "syscall" + "time" + + "github.com/nrepl/proof/internal/report" + "github.com/nrepl/proof/internal/serve" +) + +func serveFlags(out io.Writer, o *clientOptions) *flag.FlagSet { + fs := flag.NewFlagSet("serve", flag.ContinueOnError) + fs.SetOutput(out) + o.register(fs) + return fs +} + +// runServe is a server for clients to run their tests against until it's +// interrupted, and then grades everything the clients sent. +func runServe(args []string) int { + ctx, stop := signal.NotifyContext(context.Background(), os.Interrupt, syscall.SIGTERM) + defer stop() + return serveUntil(ctx, args, os.Stdout, os.Stderr, nil) +} + +// serveUntil does the work of runServe, grading the traffic once ctx is +// done. listening, if not nil, gets the address clients should use. +func serveUntil(ctx context.Context, args []string, stdout, stderr io.Writer, listening func(addr string)) int { + // The server logs to stderr from goroutines of its own. + stderr = &syncWriter{w: stderr} + var o clientOptions + fs := serveFlags(stderr, &o) + if err := fs.Parse(args); err != nil { + return 2 + } + srv, err := serve.Listen(o.listen, version, fs.Args()) + if err != nil { + fmt.Fprintln(stderr, "proof:", err) + return 2 + } + srv.Logf = logTo(stderr) + go srv.Serve() + started := time.Now() + name := "proof serve" + if fs.NArg() > 0 { + name += " (" + strings.Join(fs.Args(), ", ") + ")" + } + fmt.Fprintf(stderr, "Running %s on %s. Connect your client and press Ctrl-C when it's done.\n", name, srv.Addr()) + if listening != nil { + listening(srv.Addr()) + } + + <-ctx.Done() + fmt.Fprintln(stderr) + r := report.Run{Proof: version, Server: "client traffic to " + name, Address: srv.Addr(), Started: started} + return gradeClients(stdout, stderr, r, srv.Stop(time.Second), o, nil) +} diff --git a/cmd/proof/serve_test.go b/cmd/proof/serve_test.go new file mode 100644 index 0000000..d9b1649 --- /dev/null +++ b/cmd/proof/serve_test.go @@ -0,0 +1,34 @@ +package main + +import ( + "strings" + "testing" +) + +func TestServeExitCodes(t *testing.T) { + cases := []struct { + name string + args []string + frames []string + code int + output string + }{ + {"well-behaved client", nil, []string{"d2:id1:14:code7:(+ 1 2)2:op4:evale"}, 0, "client traffic to proof serve"}, + {"with scenarios", []string{"last-value", "no-err"}, []string{"d2:id1:12:op8:describee"}, 0, + "client traffic to proof serve (last-value, no-err)"}, + {"request without an id", nil, []string{"d2:op8:describee"}, 1, "a request without an id"}, + {"no client", nil, nil, 3, "nothing to check"}, + {"unknown scenario", []string{"no-such-scenario"}, nil, 2, `unknown scenario "no-such-scenario"`}, + } + for _, c := range cases { + t.Run(c.name, func(t *testing.T) { + code, stdout, stderr := runWithClient(t, serveUntil, c.frames, c.args...) + if code != c.code { + t.Errorf("exit status %d, want %d\nstdout:\n%s\nstderr:\n%s", code, c.code, stdout, stderr) + } + if !strings.Contains(stdout+stderr, c.output) { + t.Errorf("output doesn't mention %q\nstdout:\n%s\nstderr:\n%s", c.output, stdout, stderr) + } + }) + } +} diff --git a/doc/design.md b/doc/design.md index d3beb12..e3189b7 100644 --- a/doc/design.md +++ b/doc/design.md @@ -155,8 +155,36 @@ how they react. As a proxy can only see the wire, proof checks what a client sends and not what it does with the replies. Whether a client copes with output that arrives after `done`, or with output split into many messages, is a -different problem, which needs a server that misbehaves on purpose (see -[Future Plans](#future-plans)). +different problem, which `proof serve` takes care of. + +## Testing Clients + +`proof serve` is an nREPL server for the test suites of clients. On its +own it behaves like nREPL, and scenarios make it behave like other +servers where they differ from nREPL. The scenarios aren't made up. +Every one of them is something a server in the compatibility matrix does +(e.g. `last-value` is what Basilisp, jank and dialtone do), or something +TCP can do to any server's replies (e.g. deliver a message in pieces). +That keeps the list down to the differences a client will actually run +into, and the tests make sure it stays that way. proof runs its own +checks against every scenario, and a scenario has to get the same +verdicts as the servers it names. All the scenarios of jank together get jank's +column of the matrix, for instance. + +The grading works the other way around here. proof can't see what a +client shows its users, so the client's own tests have to do that. Most +scenarios change only the shape of the replies and not what a user +should see, so a client's tests can check the same things under every +scenario. [Autobahn|Testsuite](https://github.com/crossbario/autobahn-testsuite) +does something similar for WebSocket clients. proof still checks the +requests, though, just like `proof proxy` does. + +proof is not a Clojure implementation, so `proof serve` understands only +a small piece of Clojure. That's enough for the snippets of the nREPL +profile, the code CIDER sends when it connects and what client tests +need (output, values, errors, input, something to interrupt), and it +gives the same replies as nREPL 1.7.0 for the same code. Code in other +languages wouldn't help, as no client's tests send Erlang to a server. Some rules are about what a client leaves behind - sessions that were never closed and `need-input` that was never answered. They apply only @@ -207,9 +235,11 @@ Here's how the codebase is organized: cmd/proof the command-line interface internal/report text and JSON reports, the compatibility matrix internal/checks the checks, the wire checks, the client rules and a fake server for testing them +internal/proxy forwarding the traffic between a client and a server +internal/serve a server for client tests, which acts like other servers on request +internal/clients accepting clients and recording what they say, for proxy and serve internal/check running and grading checks, expected failures internal/server starting servers -internal/proxy forwarding and recording the traffic between a client and a server internal/profile loading profiles nrepl a client that records all messages bencode a strict bencode implementation @@ -248,8 +278,9 @@ to have a single test suite that every server can be checked against. ## Future Plans At this point proof covers the core of the protocol (`describe`, unknown -ops, sessions, `eval`, `stdin` and the wire format) and the requests of -clients. Here's what's planned next: +ops, sessions, `eval`, `stdin` and the wire format), the requests of +clients and the server differences clients have to deal with. Here's +what's planned next: - checks for `interrupt`, which every interactive client relies on - checks for `completions`, `lookup` and `load-file` (for servers that @@ -260,9 +291,6 @@ clients. Here's what's planned next: or Calva connect to a server), so a report can tell you directly whether CIDER will work with your server (`proof proxy` already sees this traffic, it just doesn't save it yet) -- a server that misbehaves on purpose (late output, output split into - many messages, unusual status values), for client test suites to run - against - more servers in the compatibility matrix and a proper home for the matrix itself - incorporating [Spec Changes](spec-changes.md) into the spec, so that diff --git a/doc/hacking.md b/doc/hacking.md index 0129d37..63bceda 100644 --- a/doc/hacking.md +++ b/doc/hacking.md @@ -44,9 +44,11 @@ $ bin/proof run profiles/clojure.toml | `nrepl` | An nREPL client that records all the messages it sends and receives. | | `internal/profile` | Loading and validating profiles. | | `internal/server` | Starting servers and figuring out their ports. | -| `internal/check` | The checks framework (`Check`, `T`, `Rule`), grading and expected failures. It doesn't know anything about specific ops. | +| `internal/check` | The checks framework (`Check`, `T`, `Rule`), grading and expected failures. It doesn't know anything about specific ops. `checktest` has helpers for testing checks. | | `internal/checks` | The checks (`describe.go`, `op.go`, `session.go` and `eval.go`), the wire checks (`wire.go`), the client rules (`client.go`), the links to client and server code (`refs.go`), the fake server used to test all of them (`fake_test.go`) and a scripted client for testing the client rules (`client_test.go`). | -| `internal/proxy` | Forwarding the traffic between a client and a server and recording it, for `proof proxy`. | +| `internal/clients` | Accepting clients, recording what they say and stopping, for `proof proxy` and `proof serve`. | +| `internal/proxy` | Forwarding the traffic between a client and a server, for `proof proxy`. | +| `internal/serve` | The server behind `proof serve`: its piece of Clojure (`lang.go` and `eval.go`), the scenarios (`scenarios.go`), sessions (`session.go`) and the server itself (`serve.go`). | | `internal/report` | Text and JSON reports and the compatibility matrix. | | `profiles` | The profiles for the servers in the compatibility matrix. | | `doc/spec-changes.md` | All the gaps and disagreements found in the draft spec. | @@ -71,6 +73,11 @@ The client rules are tested the same way. `client_test.go` has a scripted client that talks to the fake server through the proxy and can be told to make one mistake at a time (see `clientQuirks`). +`proof serve` gets checked by proof itself. `matrix_test.go` runs all the +checks against it with the snippets of `profiles/clojure.toml`, so it has +to pass everything on its own, and with each scenario it has to get the +same verdicts as the servers the scenario names. + Before submitting any changes make sure the code is formatted properly and the tests pass with the race detector enabled: @@ -218,6 +225,22 @@ to a socket and prints the replies is all you need for that. Every client rule needs a quirk in `clientQuirks` and an entry in the table in `TestClientRulesCatchMistakes` (both in `client_test.go`). +## Adding a Scenario + +The scenarios of `proof serve` live in `internal/serve/scenarios.go`. +Each one sets a field of `behavior`, which the server checks wherever it +does something differently (e.g. `noErr` in `runEval`). A scenario has to +be something a real server does, and `Who` says which ones. To find out +exactly what a server sends, connect to it and print the replies, the +same way you would for a client rule. + +If the scenario changes the verdict of some check, add it to the table in +`TestScenariosGetTheVerdictsOfTheirServers` (in `matrix_test.go`) with +the verdicts its servers get in the matrix, and to the columns of those +servers further down. Otherwise make sure it does what it says in +`TestScenarioShapes` (in `serve_test.go`). The tables in `doc/usage.md` +list all the scenarios as well. + ## Adding a Server To add a server to the compatibility matrix: diff --git a/doc/usage.md b/doc/usage.md index d8c931e..18d8f44 100644 --- a/doc/usage.md +++ b/doc/usage.md @@ -2,7 +2,7 @@ This section of the documentation covers everything you need to check an nREPL server with proof, from installing proof to running it in your -server's CI, and how to check an nREPL client with it. The details of the +server's CI, and how to check and test an nREPL client with it. The details of the profile format are covered separately in [Profiles](profiles.md). ## Installation @@ -149,8 +149,8 @@ Here are the options supported by `proof run`: | 3 | Some checks couldn't run at all. Usually this means that the server died during the run. | There are a couple of other commands as well. `proof list` shows all -the checks along with their severity and `proof version` shows the -version of proof. +the checks along with their severity (and the scenarios of +`proof serve`), and `proof version` shows the version of proof. ## Running proof in CI @@ -292,6 +292,100 @@ client fails some checks. The `kill -0` makes sure the step doesn't wait forever if proof can't start the server, and the connections `nc` makes don't count, as they don't send anything. +## Testing Your Client Against Other Servers + +The proxy checks what your client sends, but not what it does with the +replies. For that there's `proof serve`, a small nREPL server for your +client's tests. Out of the box it behaves like nREPL itself, and each +scenario you give it makes it behave like some other server in one +particular way: + +```shell +$ proof serve -listen 127.0.0.1:7888 last-value no-err +Running proof serve (last-value, no-err) on 127.0.0.1:7888. Connect your client and press Ctrl-C when it's done. +``` + +With these two scenarios it sends only the value of the last form (like +Basilisp, jank and dialtone) and drops whatever the code prints to +stderr (like Basilisp, dialtone and repartee). `proof list` shows all +the scenarios, and every one of them is something a real server does +(or something TCP can do to the replies): + +| Scenario | What changes | Who does it | +|---|---|---| +| `split-output` | Output comes one character per message | nREPL splits long output, Basilisp sends `println`'s newline on its own | +| `empty-messages` | Replies to `eval` include messages with nothing but `id` and `session` | jank | +| `last-value` | Only the value of the last form is sent | Basilisp, jank, dialtone | +| `no-err` | What the code prints to stderr never reaches the client | Basilisp, dialtone, repartee | +| `error-with-done` | `eval-error` comes in the same message as `done` | jank | +| `no-op-echo` | Replies to unknown ops don't say which op it was | ClojureCLR, Basilisp, jank, dialtone, repartee | +| `no-close-op` | `describe` doesn't list `close`, even though `close` works | jank | +| `no-interrupt` | There's no `interrupt` op | Basilisp, jank | +| `no-stdin` | There's no `stdin` op, and reading input gets an empty string right away | Basilisp | +| `string-versions` | `versions.proof` is a plain string rather than a dict | Babashka, for `versions.babashka` | +| `no-session-closed` | `close` replies with `done` alone | Basilisp, jank, dialtone, repartee | +| `shared-state` | Sessions on the same connection share `*1`, `*e` and the current namespace | ClojureCLR, Basilisp, jank | +| `socket-sessions` | A session only exists on the connection that cloned it | ClojureCLR, Basilisp, jank | +| `any-session` | Requests for sessions that don't exist run in a new session | ClojureCLR, Basilisp, jank | +| `ns-fallback` | An eval in a namespace that doesn't exist runs in the current one | jank | +| `ns-error` | An eval in a namespace that doesn't exist fails without `namespace-not-found` | Basilisp | +| `eof-error` | Reading past the end of input fails instead of returning `nil` | nREPL 1.7.0 | +| `unsorted-keys` | The keys of reply dicts aren't sorted | jank | +| `byte-writes` | Replies are written a byte at a time | any server | +| `batched-writes` | Replies are held back and written together until the eval waits or ends | any server | +| `hang-up` | The server closes the connection instead of answering an `eval` | any server that crashes | + +Most scenarios change only how the replies look on the wire, not what a +user should end up seeing. Evaluating `(println "hi") (+ 1 2)` should +show `hi` and `3` with `split-output`, `byte-writes` or +`empty-messages` just like it does without them. So the easiest way to +use `proof serve` is to write your tests the way you'd check things by +hand (e.g. "this eval shows this output and this value") and run them +once for every scenario. + +proof can't evaluate real Clojure, of course. Instead it understands just +enough of it for tests, and gives the same replies as nREPL 1.7.0 does +for the same code: + +| Code | What it does | +|---|---| +| Integers, strings, keywords, `nil`, booleans, vectors, maps and quoted forms | Evaluate to themselves | +| `(+ 1 2)`, `-`, `*`, `/`, `inc`, `dec` | Arithmetic on integers and ratios, where `(/ 1 0)` throws | +| `(str ...)`, `(apply f ... coll)`, `(repeat n x)` | e.g. `(apply str (repeat 100000 "x"))` for a really long value | +| `(print ...)`, `(println ...)`, `(pr ...)`, `(prn ...)`, `(flush)` | Output, which `(binding [*out* *err*] ...)` turns into error output | +| `(read-line)` | Asks for input with `need-input` | +| `(throw (ex-info "message" {}))` | An `eval-error`, with the same `err` and `ex` as nREPL | +| `(Thread/sleep ms)` | Something to `interrupt` | +| `(future ...)` | Runs once the eval is done, so its output arrives after `done` | +| `(def x 1)`, `x`, `#'x`, `(resolve 'x)`, `@#'x` | Definitions, which all sessions share | +| `(ns foo)`, `(in-ns 'foo)`, `*ns*` | Namespaces | +| `*1`, `*2`, `*3`, `*e` | The last results and the last exception in the session | +| `do`, `if`, `when`, `let`, `when-let` | The usual | +| `(require ...)` | Nothing, as there's nothing to load | + +Other functions get the error Clojure gives for a symbol it can't +resolve, and syntax proof doesn't read (e.g. sets or anonymous functions) +gets a read error. That's enough for CIDER to connect and work, and it's +all you need for checking output, values, errors, input and interrupts. + +When you stop it, `proof serve` checks the requests your client sent, +just like `proof proxy` does, with the same report, options (except for +`-address`) and exit codes. A test suite can run it in CI like this: + +```shell +for scenario in "" split-output last-value no-err byte-writes; do + proof serve -listen 127.0.0.1:7888 $scenario & + serve=$! + until nc -z 127.0.0.1 7888; do + kill -0 $serve || exit 2 + sleep 1 + done + # run your tests against port 7888 here + kill -INT $serve + wait $serve || exit 1 +done +``` + ## Troubleshooting This section lists the most common problems you may encounter while diff --git a/internal/check/checktest/checktest.go b/internal/check/checktest/checktest.go new file mode 100644 index 0000000..d2cdfe0 --- /dev/null +++ b/internal/check/checktest/checktest.go @@ -0,0 +1,45 @@ +// Package checktest helps test the checks and rules, and the servers they +// check. +package checktest + +import ( + "sort" + "testing" + + "github.com/nrepl/proof/internal/check" +) + +// ByID maps results to the ids of their checks or rules. +func ByID(results []check.Result) map[string]check.Result { + byID := make(map[string]check.Result, len(results)) + for _, r := range results { + byID[r.ID] = r + } + return byID +} + +// Verdicts makes sure each check or rule in want got the listed verdict, +// and that everything else passed. +func Verdicts(t testing.TB, results map[string]check.Result, want map[string]check.Verdict) { + t.Helper() + var ids []string + for id := range results { + ids = append(ids, id) + } + sort.Strings(ids) + for _, id := range ids { + r := results[id] + w, listed := want[id] + if !listed { + w = check.Pass + } + if r.Verdict != w { + t.Errorf("%s: got %s, want %s %v", id, r.Verdict, w, r.Details) + } + } + for id := range want { + if _, ok := results[id]; !ok { + t.Errorf("%s: no such check or rule", id) + } + } +} diff --git a/internal/checks/checks_test.go b/internal/checks/checks_test.go index be5ba97..3d74908 100644 --- a/internal/checks/checks_test.go +++ b/internal/checks/checks_test.go @@ -1,12 +1,12 @@ package checks import ( - "sort" "strings" "testing" "time" "github.com/nrepl/proof/internal/check" + "github.com/nrepl/proof/internal/check/checktest" "github.com/nrepl/proof/internal/profile" ) @@ -30,41 +30,11 @@ var fakeProfile = &profile.Profile{ func runFake(t *testing.T, q quirks) map[string]check.Result { t.Helper() env := &check.Env{Profile: fakeProfile, Addr: startFake(t, q), Settle: 20 * time.Millisecond} - results := map[string]check.Result{} - for _, r := range check.Run(env, All(), WireRules()) { - results[r.ID] = r - } - return results -} - -// checkVerdicts makes sure each check or rule in want got the listed -// verdict, and that everything else passed. -func checkVerdicts(t *testing.T, results map[string]check.Result, want map[string]check.Verdict) { - t.Helper() - var ids []string - for id := range results { - ids = append(ids, id) - } - sort.Strings(ids) - for _, id := range ids { - r := results[id] - w, listed := want[id] - if !listed { - w = check.Pass - } - if r.Verdict != w { - t.Errorf("%s: got %s, want %s %v", id, r.Verdict, w, r.Details) - } - } - for id := range want { - if _, ok := results[id]; !ok { - t.Errorf("%s: no such check or rule", id) - } - } + return checktest.ByID(check.Run(env, All(), WireRules())) } func TestWellBehavedServerPassesEverything(t *testing.T) { - checkVerdicts(t, runFake(t, quirks{}), nil) + checktest.Verdicts(t, runFake(t, quirks{}), nil) } // Each misbehaviour must produce exactly the listed verdicts, and every @@ -125,7 +95,7 @@ func TestChecksCatchMisbehaviour(t *testing.T) { for _, c := range cases { t.Run(c.name, func(t *testing.T) { t.Parallel() - checkVerdicts(t, runFake(t, c.q), c.want) + checktest.Verdicts(t, runFake(t, c.q), c.want) }) } } diff --git a/internal/checks/client_test.go b/internal/checks/client_test.go index 6713341..e331efe 100644 --- a/internal/checks/client_test.go +++ b/internal/checks/client_test.go @@ -8,6 +8,7 @@ import ( "github.com/nrepl/proof/bencode" "github.com/nrepl/proof/internal/check" + "github.com/nrepl/proof/internal/check/checktest" "github.com/nrepl/proof/internal/proxy" "github.com/nrepl/proof/nrepl" ) @@ -150,15 +151,11 @@ func runClient(t *testing.T, q clientQuirks) map[string]check.Result { if len(traffic) != want { t.Fatalf("got %d transcripts, want %d", len(traffic), want) } - results := map[string]check.Result{} - for _, r := range check.Grade(ClientRules(), traffic) { - results[r.ID] = r - } - return results + return checktest.ByID(check.Grade(ClientRules(), traffic)) } func TestWellBehavedClientPassesEverything(t *testing.T) { - checkVerdicts(t, runClient(t, clientQuirks{}), nil) + checktest.Verdicts(t, runClient(t, clientQuirks{}), nil) } // Each mistake must produce exactly the listed verdicts, and every other @@ -189,7 +186,7 @@ func TestClientRulesCatchMistakes(t *testing.T) { for _, c := range cases { t.Run(c.name, func(t *testing.T) { t.Parallel() - checkVerdicts(t, runClient(t, c.q), c.want) + checktest.Verdicts(t, runClient(t, c.q), c.want) }) } } @@ -241,11 +238,8 @@ func TestLeftoversCountOnlyWhenTheClientLeftFirst(t *testing.T) { ev.Time = time.Unix(int64(i), 0) events[i] = ev } - results := map[string]check.Result{} - for _, r := range check.Grade(ClientRules(), []check.Traffic{{Label: "connection 1", Events: events}}) { - results[r.ID] = r - } - checkVerdicts(t, results, map[string]check.Verdict{c.rule: c.want}) + results := checktest.ByID(check.Grade(ClientRules(), []check.Traffic{{Label: "connection 1", Events: events}})) + checktest.Verdicts(t, results, map[string]check.Verdict{c.rule: c.want}) }) } } diff --git a/internal/clients/clients.go b/internal/clients/clients.go new file mode 100644 index 0000000..1b22829 --- /dev/null +++ b/internal/clients/clients.go @@ -0,0 +1,205 @@ +// Package clients accepts the connections of nREPL clients and records +// everything said on them, so the commands that check clients can grade +// their requests afterwards. +package clients + +import ( + "errors" + "net" + "slices" + "strconv" + "sync" + "time" + + "github.com/nrepl/proof/internal/check" + "github.com/nrepl/proof/nrepl" +) + +// Listener accepts clients and numbers their connections. The transcripts +// are from the client's point of view: requests are Sent and replies are +// Received. +type Listener struct { + ln net.Listener + // Logf, if set before Serve, is told about connections coming and + // going. + Logf func(format string, args ...any) + + mu sync.Mutex + conns []*Conn + stopped bool +} + +// Conn is a client's connection. +type Conn struct { + // N numbers the connections in the order they were accepted, from 1. + N int + Client net.Conn + l *Listener + // gone is closed once the client's side of the connection is over, + // and done once its handler has returned. + gone, done chan struct{} + goneOnce sync.Once + // also are closed along with Client when proof stops. + also []net.Conn + + mu sync.Mutex + events []nrepl.Event +} + +// Listen starts accepting clients on addr. Call Serve to handle them. +func Listen(addr string) (*Listener, error) { + ln, err := net.Listen("tcp", addr) + if err != nil { + return nil, err + } + return New(ln), nil +} + +// New accepts clients from ln. +func New(ln net.Listener) *Listener { return &Listener{ln: ln} } + +// Addr is the address clients should connect to. +func (l *Listener) Addr() string { return l.ln.Addr().String() } + +// Log passes a message on to Logf. +func (l *Listener) Log(format string, args ...any) { + if l.Logf != nil { + l.Logf(format, args...) + } +} + +// Serve accepts clients until Stop is called, and handles each of them in +// a goroutine of its own. The connection is over (and closed) once handle +// returns. +func (l *Listener) Serve(handle func(*Conn)) { + var delay time.Duration + for { + nc, err := l.ln.Accept() + if errors.Is(err, net.ErrClosed) { + return + } + if err != nil { + // Most likely out of file descriptors, which passes as the + // clients hang up. + delay = min(max(2*delay, 5*time.Millisecond), time.Second) + l.Log("Couldn't accept a client: %v", err) + time.Sleep(delay) + continue + } + delay = 0 + l.mu.Lock() + if l.stopped { + l.mu.Unlock() + nc.Close() + return + } + c := &Conn{N: len(l.conns) + 1, Client: nc, l: l, gone: make(chan struct{}), done: make(chan struct{})} + l.conns = append(l.conns, c) + l.mu.Unlock() + go func() { + defer close(c.done) + defer c.ClientGone() + defer nc.Close() + handle(c) + }() + } +} + +// Stop stops accepting clients. Clients that are still connected get up +// to grace to hang up on their own (e.g. when a test suite has just +// finished), and are then disconnected. Stop returns once every handler +// has. +func (l *Listener) Stop(grace time.Duration) { + l.mu.Lock() + l.stopped = true + conns := l.conns + l.mu.Unlock() + l.ln.Close() + + deadline := time.NewTimer(grace) + defer deadline.Stop() +wait: + for _, c := range conns { + select { + case <-c.gone: + case <-deadline.C: + break wait + } + } + l.mu.Lock() + for _, c := range conns { + select { + case <-c.gone: + default: + l.Log("Connection %d was still open", c.N) + } + c.Client.Close() + for _, nc := range c.also { + nc.Close() + } + } + l.mu.Unlock() + for _, c := range conns { + <-c.done + } +} + +// Conns returns every connection so far, in the order they were opened. +func (l *Listener) Conns() []*Conn { + l.mu.Lock() + defer l.mu.Unlock() + return slices.Clone(l.conns) +} + +// Traffic returns the transcript of every connection in which the client +// sent something, labelled with the numbers the log uses. Connections that +// only probed the port (e.g. a CI script waiting for proof to start) have +// nothing to check. +func (l *Listener) Traffic() []check.Traffic { + var traffic []check.Traffic + for _, c := range l.Conns() { + events := c.Events() + if slices.ContainsFunc(events, func(ev nrepl.Event) bool { return ev.Dir == nrepl.Sent && !ev.Closed }) { + traffic = append(traffic, check.Traffic{Label: "connection " + strconv.Itoa(c.N), Events: events}) + } + } + return traffic +} + +// Record adds an event to the connection's transcript. +func (c *Conn) Record(ev nrepl.Event) { + c.mu.Lock() + c.events = append(c.events, ev) + c.mu.Unlock() +} + +// Events returns the transcript so far. +func (c *Conn) Events() []nrepl.Event { + c.mu.Lock() + defer c.mu.Unlock() + return slices.Clone(c.events) +} + +// ClientGone marks the client's side of the connection as over, while the +// handler may still be busy with the rest (e.g. with a server that's slow +// to hang up). Serve does this anyway once the handler returns. +func (c *Conn) ClientGone() { c.goneOnce.Do(func() { close(c.gone) }) } + +// Gone is closed once the client's side of the connection is over. +func (c *Conn) Gone() <-chan struct{} { return c.gone } + +// Done is closed once the connection is over. +func (c *Conn) Done() <-chan struct{} { return c.done } + +// Also makes Stop close nc along with the client's connection. If proof is +// stopping already, it closes nc right away and reports false. +func (c *Conn) Also(nc net.Conn) bool { + c.l.mu.Lock() + defer c.l.mu.Unlock() + if c.l.stopped { + nc.Close() + return false + } + c.also = append(c.also, nc) + return true +} diff --git a/internal/clients/clients_test.go b/internal/clients/clients_test.go new file mode 100644 index 0000000..43559c0 --- /dev/null +++ b/internal/clients/clients_test.go @@ -0,0 +1,194 @@ +package clients + +import ( + "bytes" + "errors" + "io" + "net" + "slices" + "strings" + "sync" + "testing" + "time" + + "github.com/nrepl/proof/nrepl" +) + +// echo records what a client sends and sends it back, until the client +// hangs up. +func echo(c *Conn) { + buf := make([]byte, 1024) + for { + n, err := c.Client.Read(buf) + if n > 0 { + c.Record(nrepl.Event{Dir: nrepl.Sent, Raw: slices.Clone(buf[:n])}) + c.Client.Write(buf[:n]) + } + if err != nil { + if !errors.Is(err, net.ErrClosed) { + c.Record(nrepl.Event{Dir: nrepl.Sent, Closed: true}) + } + return + } + } +} + +// start serves clients from ln with handle, and returns what it logged so +// far. +func start(t *testing.T, ln net.Listener, handle func(*Conn)) (*Listener, func() string) { + t.Helper() + l := New(ln) + var logged bytes.Buffer + var mu sync.Mutex + l.Logf = func(format string, args ...any) { + mu.Lock() + defer mu.Unlock() + logged.WriteString(format + "\n") + } + go l.Serve(handle) + t.Cleanup(func() { l.Stop(0) }) + return l, func() string { + mu.Lock() + defer mu.Unlock() + return logged.String() + } +} + +func listen(t *testing.T) net.Listener { + t.Helper() + ln, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + return ln +} + +func dial(t *testing.T, l *Listener) net.Conn { + t.Helper() + nc, err := net.Dial("tcp", l.Addr()) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { nc.Close() }) + return nc +} + +// sendAndHangUp sends data and waits for the echo, which also means the +// listener has it. +func sendAndHangUp(t *testing.T, l *Listener, data string) { + t.Helper() + nc := dial(t, l) + nc.Write([]byte(data)) + nc.(*net.TCPConn).CloseWrite() + nc.SetReadDeadline(time.Now().Add(2 * time.Second)) + io.ReadAll(nc) +} + +// Connections that never send anything (e.g. a script checking whether the +// port is open yet) aren't clients, but they keep their number so the +// report matches the log. +func TestConnectionsThatSendNothingAreLeftOut(t *testing.T) { + l, _ := start(t, listen(t), echo) + for _, data := range []string{"", "d2:op8:describee"} { + sendAndHangUp(t, l, data) + } + l.Stop(time.Second) + if traffic := l.Traffic(); len(traffic) != 1 || traffic[0].Label != "connection 2" { + t.Errorf("got %v, want just connection 2", traffic) + } +} + +func TestStopDisconnectsClientsStillConnected(t *testing.T) { + l, logged := start(t, listen(t), echo) + sendAndHangUp(t, l, "1") + still := dial(t, l) + still.Write([]byte("2")) + still.Read(make([]byte, 1)) + + started := time.Now() + l.Stop(50 * time.Millisecond) + if d := time.Since(started); d > time.Second { + t.Errorf("Stop took %s", d) + } + still.SetReadDeadline(time.Now().Add(time.Second)) + if _, err := still.Read(make([]byte, 1)); err != io.EOF { + t.Errorf("client read: %v, want EOF", err) + } + if log := logged(); log != "Connection %d was still open\n" { + t.Errorf("only connection 2 should've been still open:\n%s", log) + } +} + +// A handler can say the client is gone before it's done itself, and Stop +// closes whatever else the connection needs along with it. +func TestStopClosesWhatTheConnectionHolds(t *testing.T) { + held, other := net.Pipe() + defer other.Close() + l, logged := start(t, listen(t), func(c *Conn) { + c.Also(held) + c.ClientGone() + // Until Stop closes it. + io.Copy(io.Discard, held) + }) + dial(t, l).Write([]byte("1")) + for len(l.Conns()) == 0 { + time.Sleep(time.Millisecond) + } + <-l.Conns()[0].Gone() + + started := time.Now() + l.Stop(5 * time.Second) + if d := time.Since(started); d > time.Second { + t.Errorf("Stop waited %s for a client that was gone", d) + } + if strings.Contains(logged(), "still open") { + t.Errorf("log:\n%s", logged()) + } + held2, other2 := net.Pipe() + defer other2.Close() + if l.Conns()[0].Also(held2) { + t.Error("Also took a connection after Stop") + } + if _, err := held2.Write([]byte("x")); !errors.Is(err, io.ErrClosedPipe) { + t.Errorf("Also should've closed the connection it didn't take: %v", err) + } +} + +// The connection is over once its handler returns, so the client hears +// about it right away. +func TestClosesConnectionsOnceHandled(t *testing.T) { + l, _ := start(t, listen(t), func(c *Conn) {}) + nc := dial(t, l) + nc.SetReadDeadline(time.Now().Add(time.Second)) + if _, err := nc.Read(make([]byte, 1)); err != io.EOF { + t.Errorf("client read: %v, want EOF", err) + } +} + +// failOnce is a listener whose first Accept fails, the way it does when +// proof runs out of file descriptors. +type failOnce struct { + net.Listener + failed bool +} + +func (l *failOnce) Accept() (net.Conn, error) { + if !l.failed { + l.failed = true + return nil, errors.New("too many open files") + } + return l.Listener.Accept() +} + +func TestKeepsAcceptingAfterAnError(t *testing.T) { + l, logged := start(t, &failOnce{Listener: listen(t)}, echo) + nc := dial(t, l) + nc.Write([]byte("1")) + nc.SetReadDeadline(time.Now().Add(2 * time.Second)) + if _, err := nc.Read(make([]byte, 1)); err != nil { + t.Errorf("no reply: %v", err) + } + if !strings.Contains(logged(), "Couldn't accept a client") { + t.Errorf("the error wasn't logged:\n%s", logged()) + } +} diff --git a/internal/proxy/proxy.go b/internal/proxy/proxy.go index 4d7adc6..634b2c4 100644 --- a/internal/proxy/proxy.go +++ b/internal/proxy/proxy.go @@ -8,143 +8,88 @@ import ( "errors" "io" "net" - "strconv" - "sync" "time" "github.com/nrepl/proof/bencode" "github.com/nrepl/proof/internal/check" + "github.com/nrepl/proof/internal/clients" "github.com/nrepl/proof/nrepl" ) // Proxy accepts clients and connects each of them to the upstream server. -// The transcripts are from the client's point of view: requests are Sent -// and replies are Received. type Proxy struct { + *clients.Listener upstream string - ln net.Listener - // Logf, if set before Serve, is told about connections coming and - // going. - Logf func(format string, args ...any) - - mu sync.Mutex - conns []*conn - stopped bool } -type conn struct { - n int - client, server net.Conn - // clientDone is closed once the client's side of the connection has - // ended (or never got going). - clientDone chan struct{} - // done is closed once the connection is over and both sockets are - // closed. - done chan struct{} - - mu sync.Mutex - events []nrepl.Event -} +// conn is a client's connection, which the relays record. +type conn struct{ *clients.Conn } // Listen starts accepting clients on addr. Call Serve to handle them. func Listen(addr, upstream string) (*Proxy, error) { - ln, err := net.Listen("tcp", addr) + l, err := clients.Listen(addr) if err != nil { return nil, err } - return &Proxy{upstream: upstream, ln: ln}, nil -} - -// Addr is the address clients should connect to. -func (p *Proxy) Addr() string { return p.ln.Addr().String() } - -func (p *Proxy) logf(format string, args ...any) { - if p.Logf != nil { - p.Logf(format, args...) - } + return &Proxy{Listener: l, upstream: upstream}, nil } // Serve accepts clients until Stop is called. -func (p *Proxy) Serve() { - var delay time.Duration - for { - nc, err := p.ln.Accept() - if errors.Is(err, net.ErrClosed) { - return - } - if err != nil { - // Most likely out of file descriptors, which passes as the - // clients hang up. - delay = min(max(2*delay, 5*time.Millisecond), time.Second) - p.logf("Couldn't accept a client: %v", err) - time.Sleep(delay) - continue - } - delay = 0 - p.mu.Lock() - if p.stopped { - p.mu.Unlock() - nc.Close() - return - } - c := &conn{n: len(p.conns) + 1, client: nc, clientDone: make(chan struct{}), done: make(chan struct{})} - p.conns = append(p.conns, c) - p.mu.Unlock() - go p.handle(c) - } +func (p *Proxy) Serve() { p.Listener.Serve(p.handle) } + +// Stop stops accepting clients and returns the transcript of every +// connection in which the client sent something. Clients that are still +// connected get up to grace to hang up on their own, and are then +// disconnected. +func (p *Proxy) Stop(grace time.Duration) []check.Traffic { + p.Listener.Stop(grace) + return p.Traffic() } -func (p *Proxy) handle(c *conn) { - defer close(c.done) +func (p *Proxy) handle(cc *clients.Conn) { server, err := net.DialTimeout("tcp", p.upstream, 5*time.Second) - p.mu.Lock() - if err == nil && p.stopped { - server.Close() + if err == nil && !cc.Also(server) { err = errors.New("proof is stopping") } - if err == nil { - c.server = server - } - p.mu.Unlock() if err != nil { - p.logf("Connection %d: couldn't connect to %s: %v", c.n, p.upstream, err) - c.client.Close() - close(c.clientDone) + p.Log("Connection %d: couldn't connect to %s: %v", cc.N, p.upstream, err) + cc.Client.Close() return } - p.logf("Connection %d opened", c.n) + c := conn{cc} + p.Log("Connection %d opened", cc.N) go func() { - defer close(c.clientDone) - end := c.relay(c.client, c.server, nrepl.Sent) + defer cc.ClientGone() + end := c.relay(cc.Client, server, nrepl.Sent) if end == dstGone { // The server is gone, but if the client hangs up before proof // cuts it off, that still counts. - end = c.passOn(c.client, io.Discard, nrepl.Sent) + end = c.passOn(cc.Client, io.Discard, nrepl.Sent) } // Some servers never hang up their end, so this is when the // connection is over as far as the client is concerned. if end == hungUp { - p.logf("Connection %d closed", c.n) + p.Log("Connection %d closed", cc.N) // Let the server see the end of the requests, while any // replies still on their way can reach the client. - closeWrite(c.server) + closeWrite(server) } }() - switch c.relay(c.server, c.client, nrepl.Received) { + switch c.relay(server, cc.Client, nrepl.Received) { case hungUp: - p.logf("Connection %d closed by the server", c.n) + p.Log("Connection %d closed by the server", cc.N) case dstGone: // The client is gone, and its side of the relay records that. // The server finds out the way it would without proof in between. - c.server.Close() - <-c.clientDone + server.Close() + <-cc.Gone() } // Without a server there's nothing more for the client to do, and // nothing it still sends can get anywhere. (Or proof is stopping.) - c.client.Close() - c.server.Close() - <-c.clientDone + cc.Client.Close() + server.Close() + <-cc.Gone() } // ending says how a relay ended. @@ -160,14 +105,14 @@ const ( // each of them (and src hanging up), and says how it ended. After a frame // that can't be decoded the rest of the stream is passed on without being // recorded, as there's no telling where the next frame starts. -func (c *conn) relay(src net.Conn, dst io.Writer, dir nrepl.Direction) ending { +func (c conn) relay(src net.Conn, dst io.Writer, dir nrepl.Direction) ending { br := bufio.NewReader(src) dec := bencode.NewDecoder(br) for { v, err := dec.Decode() ev, isFrame := nrepl.DecodeEvent(dir, v, err) if isFrame { - c.record(ev) + c.Record(ev) if len(v.Raw) > 0 { if _, werr := dst.Write(v.Raw); werr != nil { return dstGone @@ -185,7 +130,7 @@ func (c *conn) relay(src net.Conn, dst io.Writer, dir nrepl.Direction) ending { } // passOn copies the rest of src to dst as it is, and says how that ended. -func (c *conn) passOn(src io.Reader, dst io.Writer, dir nrepl.Direction) ending { +func (c conn) passOn(src io.Reader, dst io.Writer, dir nrepl.Direction) ending { buf := make([]byte, 32<<10) for { n, err := src.Read(buf) @@ -201,86 +146,18 @@ func (c *conn) passOn(src io.Reader, dst io.Writer, dir nrepl.Direction) ending } // ended records src hanging up, unless it was proof that closed it. -func (c *conn) ended(dir nrepl.Direction, err error) ending { +func (c conn) ended(dir nrepl.Direction, err error) ending { if errors.Is(err, net.ErrClosed) { return stopped } // The end of the stream, or a reset (e.g. from a client that closed // its socket with replies it hadn't read yet). - c.record(nrepl.Event{Dir: dir, Time: time.Now(), Closed: true}) + c.Record(nrepl.Event{Dir: dir, Time: time.Now(), Closed: true}) return hungUp } -func (c *conn) record(ev nrepl.Event) { - c.mu.Lock() - c.events = append(c.events, ev) - c.mu.Unlock() -} - func closeWrite(nc net.Conn) { if tc, ok := nc.(*net.TCPConn); ok { tc.CloseWrite() } } - -// Stop stops accepting clients and returns the transcript of every -// connection in which the client sent something, in the order they were -// opened and labelled with the numbers the log uses. Clients that are -// still connected get up to grace to hang up on their own (e.g. when a -// test suite has just finished), and are then disconnected. -func (p *Proxy) Stop(grace time.Duration) []check.Traffic { - p.mu.Lock() - p.stopped = true - conns := p.conns - p.mu.Unlock() - p.ln.Close() - - deadline := time.NewTimer(grace) - defer deadline.Stop() -wait: - for _, c := range conns { - select { - case <-c.clientDone: - case <-deadline.C: - break wait - } - } - p.mu.Lock() - for _, c := range conns { - select { - case <-c.clientDone: - default: - p.logf("Connection %d was still open", c.n) - } - c.client.Close() - if c.server != nil { - c.server.Close() - } - } - p.mu.Unlock() - for _, c := range conns { - <-c.done - } - - var traffic []check.Traffic - for _, c := range conns { - c.mu.Lock() - if sentAnything(c.events) { - traffic = append(traffic, check.Traffic{Label: "connection " + strconv.Itoa(c.n), Events: c.events}) - } - c.mu.Unlock() - } - return traffic -} - -// sentAnything reports whether the client sent anything. Connections that -// only probed the port (e.g. a CI script waiting for proof to start) have -// nothing to check. -func sentAnything(events []nrepl.Event) bool { - for _, ev := range events { - if ev.Dir == nrepl.Sent && !ev.Closed { - return true - } - } - return false -} diff --git a/internal/proxy/proxy_test.go b/internal/proxy/proxy_test.go index 5741c73..e01724f 100644 --- a/internal/proxy/proxy_test.go +++ b/internal/proxy/proxy_test.go @@ -12,6 +12,7 @@ import ( "time" "github.com/nrepl/proof/internal/check" + "github.com/nrepl/proof/internal/clients" "github.com/nrepl/proof/nrepl" ) @@ -79,14 +80,11 @@ func within(t *testing.T, ch <-chan struct{}, what string) { } // firstConn returns the first connection p accepted. -func firstConn(t *testing.T, p *Proxy) *conn { +func firstConn(t *testing.T, p *Proxy) *clients.Conn { t.Helper() deadline := time.Now().Add(5 * time.Second) for { - p.mu.Lock() - conns := p.conns - p.mu.Unlock() - if len(conns) > 0 { + if conns := p.Conns(); len(conns) > 0 { return conns[0] } if time.Now().After(deadline) { @@ -246,14 +244,13 @@ func TestBothSocketsAreClosedOnceTheClientIsGone(t *testing.T) { c := firstConn(t, p) // Stop would close the sockets itself, so the connection has to end // on its own. - within(t, c.done, "the connection to end") + within(t, c.Done(), "the connection to end") + // The server's writes fail once proof closes its socket. within(t, writeFailed, "the server's writes to fail") - for name, nc := range map[string]net.Conn{"client": c.client, "server": c.server} { - if err := nc.SetDeadline(time.Time{}); !errors.Is(err, net.ErrClosed) { - t.Errorf("the %s socket is still open", name) - } + if err := c.Client.SetDeadline(time.Time{}); !errors.Is(err, net.ErrClosed) { + t.Error("the client's socket is still open") } - if got, want := describe(c.events), "-> {id \"1\", op \"describe\"}\n-> closed"; !strings.HasPrefix(got, want) { + if got, want := describe(c.Events()), "-> {id \"1\", op \"describe\"}\n-> closed"; !strings.HasPrefix(got, want) { t.Errorf("transcript:\n%s\nwant it to start with:\n%s", got, want) } } @@ -287,9 +284,9 @@ func TestClientHangingUpWhileItsRequestsAreOnTheirWay(t *testing.T) { nc.Close() c := firstConn(t, p) - within(t, c.done, "the connection to end") - if !slices.ContainsFunc(c.events, func(ev nrepl.Event) bool { return ev.Dir == nrepl.Sent && ev.Closed }) { - t.Errorf("no hang-up in the transcript:\n%s", describe(c.events[1:])) + within(t, c.Done(), "the connection to end") + if !slices.ContainsFunc(c.Events(), func(ev nrepl.Event) bool { return ev.Dir == nrepl.Sent && ev.Closed }) { + t.Errorf("no hang-up in the transcript:\n%s", describe(c.Events()[1:])) } } @@ -400,20 +397,6 @@ func TestStopDisconnectsClientsStillConnected(t *testing.T) { } } -// Connections that never send anything (e.g. a script checking whether the -// port is open yet) aren't clients, but they keep their number so the -// report matches the log. -func TestConnectionsThatSendNothingAreLeftOut(t *testing.T) { - p, _ := start(t, upstream(t, func(c net.Conn) { io.Copy(io.Discard, c) })) - for _, frame := range []string{"", "d2:op8:describee"} { - sendAndHangUp(t, p, frame) - } - traffic := p.Stop(time.Second) - if len(traffic) != 1 || traffic[0].Label != "connection 2" { - t.Errorf("got %v, want just connection 2", traffic) - } -} - func TestUnreachableServer(t *testing.T) { ln, err := net.Listen("tcp", "127.0.0.1:0") if err != nil { @@ -437,50 +420,6 @@ func TestUnreachableServer(t *testing.T) { } } -// failOnce is a listener whose first Accept fails, the way it does when -// proof runs out of file descriptors. -type failOnce struct { - net.Listener - failed bool -} - -func (l *failOnce) Accept() (net.Conn, error) { - if !l.failed { - l.failed = true - return nil, errors.New("too many open files") - } - return l.Listener.Accept() -} - -func TestKeepsAcceptingAfterAnError(t *testing.T) { - addr := upstream(t, func(c net.Conn) { - c.Read(make([]byte, 64)) - c.Write([]byte("d2:id1:16:statusl4:doneee")) - }) - ln, err := net.Listen("tcp", "127.0.0.1:0") - if err != nil { - t.Fatal(err) - } - p := &Proxy{upstream: addr, ln: &failOnce{Listener: ln}} - logged := logTo(p) - go p.Serve() - t.Cleanup(func() { p.Stop(0) }) - - c, err := net.Dial("tcp", p.Addr()) - if err != nil { - t.Fatal(err) - } - defer c.Close() - c.Write([]byte("d2:id1:12:op8:describee")) - c.SetReadDeadline(time.Now().Add(2 * time.Second)) - if _, err := c.Read(make([]byte, 1)); err != nil { - t.Errorf("no reply: %v", err) - } - if !strings.Contains(logged(), "Couldn't accept a client") { - t.Errorf("the error wasn't logged:\n%s", logged()) - } -} - type failing struct{} func (failing) Write(p []byte) (int, error) { return 0, errors.New("the other side is gone") } @@ -491,12 +430,12 @@ func TestFailedWritesAreNotHangUps(t *testing.T) { for _, data := range []string{"d2:op8:describee", "hello"} { client, proxied := net.Pipe() go client.Write([]byte(data)) - c := &conn{} + c := conn{&clients.Conn{}} if end := c.relay(proxied, failing{}, nrepl.Sent); end != dstGone { t.Errorf("%q: ended with %d, want dstGone", data, end) } - if strings.Contains(describe(c.events), "closed") { - t.Errorf("%q: transcript:\n%s", data, describe(c.events)) + if strings.Contains(describe(c.Events()), "closed") { + t.Errorf("%q: transcript:\n%s", data, describe(c.Events())) } client.Close() proxied.Close() diff --git a/internal/serve/eval.go b/internal/serve/eval.go new file mode 100644 index 0000000..daec347 --- /dev/null +++ b/internal/serve/eval.go @@ -0,0 +1,717 @@ +package serve + +import ( + "context" + "errors" + "fmt" + "maps" + "math" + "math/big" + "strings" + "sync" + "time" +) + +// bindings are what nREPL keeps for each session: the current namespace +// and the last results. Definitions belong to the server, as namespaces do +// in Clojure. +type bindings struct { + mu sync.Mutex + ns string + results [3]any // *1, *2 and *3 + lastErr any // *e +} + +func newBindings() *bindings { return &bindings{ns: "user"} } + +func (b *bindings) copy() *bindings { + b.mu.Lock() + defer b.mu.Unlock() + return &bindings{ns: b.ns, results: b.results, lastErr: b.lastErr} +} + +func (b *bindings) currentNS() string { + b.mu.Lock() + defer b.mu.Unlock() + return b.ns +} + +func (b *bindings) setNS(ns string) { + b.mu.Lock() + defer b.mu.Unlock() + b.ns = ns +} + +// push makes v the latest result, *1. +func (b *bindings) push(v any) { + b.mu.Lock() + defer b.mu.Unlock() + b.results = [3]any{v, b.results[0], b.results[1]} +} + +func (b *bindings) failed(ex *exception) { + b.mu.Lock() + defer b.mu.Unlock() + b.lastErr = ex +} + +// recall returns *1, *2, *3 or *e. +func (b *bindings) recall(name string) any { + b.mu.Lock() + defer b.mu.Unlock() + if name == "*e" { + return b.lastErr + } + return b.results[name[1]-'1'] +} + +// definitions are the namespaces and the vars defined in them, which all +// sessions share. The vars include clojure.core's functions. +type definitions struct { + mu sync.Mutex + nss map[string]bool + vars map[string]any // by qualified name +} + +func newDefinitions() *definitions { + d := &definitions{nss: map[string]bool{"user": true, "clojure.core": true}, vars: map[string]any{}} + for name := range builtins { + if qualified(name) { + d.vars[name] = &function{name} + } else { + d.vars["clojure.core/"+name] = &function{name} + } + } + return d +} + +func (d *definitions) lookup(qualified string) (any, bool) { + d.mu.Lock() + defer d.mu.Unlock() + v, ok := d.vars[qualified] + return v, ok +} + +func (d *definitions) define(qualified string, v any) { + d.mu.Lock() + defer d.mu.Unlock() + d.vars[qualified] = v +} + +func (d *definitions) hasNS(ns string) bool { + d.mu.Lock() + defer d.mu.Unlock() + return d.nss[ns] +} + +func (d *definitions) addNS(ns string) { + d.mu.Lock() + defer d.mu.Unlock() + d.nss[ns] = true +} + +// evaluation evaluates the forms of one eval request. +type evaluation struct { + defs *definitions + b *bindings + ns string // the namespace the code runs in + // ctx ends when proof stops, and interrupts takes interrupts. + ctx context.Context + interrupts <-chan struct{} + // waiting is called before the code waits for something. + waiting func() + // output passes on what the code prints, as "out" or "err". + output func(key, text string) + // readLine reads a line of input, or returns nil at the end of it. + readLine func(interrupts <-chan struct{}) (any, error) + // locals are the names let and when-let bind. + locals map[string]any + // futures are the futures to run once the eval is done. + futures []pending + toErr bool +} + +// pending is the body of a future, with the evaluation it runs in. +type pending struct { + e *evaluation + body []any +} + +var errStopping = errors.New("proof is stopping") + +type phase int + +const ( + running phase = iota + reading + compiling +) + +// thrown is an exception from running code, or from reading or compiling +// it. +type thrown struct { + ex *exception + phase phase + // line and col say where reading failed, and form which special form + // couldn't be compiled. + line, col int + form string +} + +func (t *thrown) Error() string { return t.ex.msg } + +func throwf(class, format string, args ...any) error { + return &thrown{ex: &exception{class: class, msg: fmt.Sprintf(format, args...)}} +} + +func compileError(form, format string, args ...any) error { + return &thrown{ex: &exception{class: "clojure.lang.Compiler$CompilerException", msg: fmt.Sprintf(format, args...)}, + phase: compiling, form: form} +} + +// errorReport is what nREPL 1.7.0 sends about an exception: the text for +// err, and the classes for ex and root-ex. +func errorReport(t *thrown, ns string) (text, ex, rootEx string) { + class := "class " + t.ex.class + switch t.phase { + case reading: + return fmt.Sprintf("Syntax error reading source at (REPL:%d:%d).\n%s\n", t.line, t.col, t.ex.msg), + class, "class java.lang.RuntimeException" + case compiling: + form := "" + if t.form != "" { + form = t.form + " " + } + return fmt.Sprintf("Syntax error compiling %sat (REPL:0:0).\n%s\n", form, t.ex.msg), class, class + } + simple := t.ex.class[strings.LastIndexByte(t.ex.class, '.')+1:] + return fmt.Sprintf("Execution error (%s) at %s/eval1 (REPL:1).\n%s\n", simple, ns, t.ex.msg), class, class +} + +func (e *evaluation) eval(form any) (any, error) { + switch f := form.(type) { + case symbol: + return e.resolve(string(f)) + case list: + if len(f) == 0 { + return f, nil + } + return e.call(f) + case vector: + xs, err := e.evalEach(f) + return vector(xs), err + case mapForm: + xs, err := e.evalEach(f) + return mapForm(xs), err + } + return form, nil +} + +func (e *evaluation) evalEach(forms []any) ([]any, error) { + vals := make([]any, 0, len(forms)) + for _, f := range forms { + v, err := e.eval(f) + if err != nil { + return nil, err + } + vals = append(vals, v) + } + return vals, nil +} + +// do evaluates forms in order and returns the value of the last one. +func (e *evaluation) do(forms []any) (any, error) { + var v any + for _, f := range forms { + var err error + if v, err = e.eval(f); err != nil { + return nil, err + } + } + return v, nil +} + +func (e *evaluation) resolve(name string) (any, error) { + if v, ok := e.locals[name]; ok { + return v, nil + } + switch name { + case "*1", "*2", "*3", "*e": + return e.b.recall(name), nil + case "*ns*": + return namespace(e.ns), nil + } + if v, ok := e.defs.lookup(e.qualify(name)); ok { + return v, nil + } + return nil, compileError("", "Unable to resolve symbol: %s in this context", name) +} + +// qualify names the var a symbol refers to: clojure.core's if it has one +// by that name, and otherwise one in the current namespace. +func (e *evaluation) qualify(name string) string { + switch { + case qualified(name): + return name + case builtins[name] != nil: + return "clojure.core/" + name + } + return e.ns + "/" + name +} + +// qualified reports whether a symbol names its namespace (like +// clojure.core/str, but unlike /). +func qualified(name string) bool { return strings.Contains(name[1:], "/") } + +// varOf returns the var a symbol refers to, or nil if there's none. +func (e *evaluation) varOf(name string) *variable { + qualified := e.qualify(name) + if _, ok := e.defs.lookup(qualified); !ok { + return nil + } + ns, n, _ := strings.Cut(qualified, "/") + return &variable{ns, n} +} + +func (e *evaluation) call(l list) (any, error) { + if s, ok := l[0].(symbol); ok && specials[string(s)] != nil { + return specials[string(s)](e, l[1:]) + } + head, err := e.eval(l[0]) + if err != nil { + return nil, err + } + args, err := e.evalEach(l[1:]) + if err != nil { + return nil, err + } + return e.apply(head, args) +} + +func (e *evaluation) apply(f any, args []any) (any, error) { + fn, ok := f.(*function) + if !ok { + return nil, castError(f, "clojure.lang.IFn") + } + return builtins[fn.name](e, args) +} + +func truthy(v any) bool { return v != nil && v != false } + +func firstSymbol(args []any) (symbol, bool) { + if len(args) == 0 { + return "", false + } + s, ok := args[0].(symbol) + return s, ok +} + +// enter switches to a namespace, creating it if needed. +func (e *evaluation) enter(ns string) { + e.defs.addNS(ns) + e.ns = ns +} + +// builtin is a special form or a function. Special forms get their +// arguments as they were read, and functions get them evaluated. +type builtin func(e *evaluation, args []any) (any, error) + +// specials are the special forms (and macros) proof serve knows, and +// builtins the functions. They're filled in by init, as some of them call +// back into the evaluator. +var specials, builtins map[string]builtin + +func init() { + specials = map[string]builtin{ + "quote": func(e *evaluation, args []any) (any, error) { + if len(args) != 1 { + return nil, arity("quote", len(args)) + } + return args[0], nil + }, + "var": func(e *evaluation, args []any) (any, error) { + s, ok := firstSymbol(args) + if !ok || len(args) != 1 { + return nil, arity("var", len(args)) + } + if v := e.varOf(string(s)); v != nil { + return v, nil + } + return nil, compileError("var", "Unable to resolve var: %s in this context", s) + }, + "do": (*evaluation).do, + "when": func(e *evaluation, args []any) (any, error) { + if len(args) == 0 { + return nil, arity("when", 0) + } + test, err := e.eval(args[0]) + if err != nil || !truthy(test) { + return nil, err + } + return e.do(args[1:]) + }, + "let": binder("let"), + "when-let": binder("when-let"), + "if": func(e *evaluation, args []any) (any, error) { + if len(args) < 2 || len(args) > 3 { + return nil, arity("if", len(args)) + } + test, err := e.eval(args[0]) + switch { + case err != nil: + return nil, err + case truthy(test): + return e.eval(args[1]) + case len(args) == 3: + return e.eval(args[2]) + } + return nil, nil + }, + "def": func(e *evaluation, args []any) (any, error) { + s, ok := firstSymbol(args) + if !ok || len(args) > 2 || strings.Contains(string(s), "/") { + return nil, throwf("java.lang.RuntimeException", "proof serve only supports (def name value)") + } + v, err := e.do(args[1:]) + if err != nil { + return nil, err + } + e.defs.define(e.ns+"/"+string(s), v) + return &variable{e.ns, string(s)}, nil + }, + "ns": func(e *evaluation, args []any) (any, error) { + s, ok := firstSymbol(args) + if !ok { + return nil, throwf("java.lang.IllegalArgumentException", "ns needs a name") + } + e.enter(string(s)) + return nil, nil + }, + "binding": func(e *evaluation, args []any) (any, error) { + if len(args) == 0 || pr(args[0]) != "[*out* *err*]" { + return nil, throwf("java.lang.RuntimeException", "proof serve only supports (binding [*out* *err*] ...)") + } + was := e.toErr + e.toErr = true + defer func() { e.toErr = was }() + return e.do(args[1:]) + }, + "throw": func(e *evaluation, args []any) (any, error) { + if len(args) != 1 { + return nil, arity("throw", len(args)) + } + v, err := e.eval(args[0]) + if err != nil { + return nil, err + } + ex, ok := v.(*exception) + if !ok { + return nil, castError(v, "java.lang.Throwable") + } + return nil, &thrown{ex: ex} + }, + "future": func(e *evaluation, args []any) (any, error) { + // It runs with the bindings and locals it was made with. + f := *e + f.futures, f.locals = nil, maps.Clone(e.locals) + e.futures = append(e.futures, pending{&f, args}) + return &future{}, nil + }, + } + + builtins = map[string]builtin{ + "+": arithmetic("+", (*big.Rat).Add, 0), + "-": arithmetic("-", (*big.Rat).Sub, 0), + "*": arithmetic("*", (*big.Rat).Mul, 1), + "/": arithmetic("/", (*big.Rat).Quo, 1), + "inc": func(e *evaluation, args []any) (any, error) { return oneMore(e, "inc", "+", args) }, + "dec": func(e *evaluation, args []any) (any, error) { return oneMore(e, "dec", "-", args) }, + "str": func(e *evaluation, args []any) (any, error) { + // Like toString in Clojure: strings as they are, namespaces by + // name and everything else as pr prints it. + var b strings.Builder + for _, a := range args { + switch x := a.(type) { + case nil: + case string: + b.WriteString(x) + case namespace: + b.WriteString(string(x)) + default: + b.WriteString(pr(x)) + } + } + return b.String(), nil + }, + "print": printer(false, false), + "println": printer(false, true), + "pr": printer(true, false), + "prn": printer(true, true), + // Output goes out as soon as it's printed. + "flush": func(e *evaluation, args []any) (any, error) { return nil, nil }, + "read-line": func(e *evaluation, args []any) (any, error) { + if len(args) != 0 { + return nil, arity("read-line", len(args)) + } + return e.readLine(e.interrupts) + }, + "ex-info": func(e *evaluation, args []any) (any, error) { + if len(args) < 2 || len(args) > 3 { + return nil, arity("ex-info", len(args)) + } + msg, ok := args[0].(string) + if !ok { + return nil, castError(args[0], "java.lang.String") + } + return &exception{class: "clojure.lang.ExceptionInfo", msg: msg, data: args[1]}, nil + }, + "ex-data": func(e *evaluation, args []any) (any, error) { + if len(args) != 1 { + return nil, arity("ex-data", len(args)) + } + if ex, ok := args[0].(*exception); ok { + return ex.data, nil + } + return nil, nil + }, + "in-ns": func(e *evaluation, args []any) (any, error) { + if len(args) != 1 { + return nil, arity("in-ns", len(args)) + } + s, ok := args[0].(symbol) + if !ok { + return nil, castError(args[0], "clojure.lang.Symbol") + } + e.enter(string(s)) + return namespace(s), nil + }, + // There's nothing to load, so requiring anything works. + "require": func(e *evaluation, args []any) (any, error) { return nil, nil }, + "resolve": func(e *evaluation, args []any) (any, error) { + if len(args) != 1 { + return nil, arity("resolve", len(args)) + } + s, ok := args[0].(symbol) + if !ok { + return nil, castError(args[0], "clojure.lang.Symbol") + } + if v := e.varOf(string(s)); v != nil { + return v, nil + } + return nil, nil + }, + "deref": func(e *evaluation, args []any) (any, error) { + if len(args) != 1 { + return nil, arity("deref", len(args)) + } + switch x := args[0].(type) { + case *variable: + v, _ := e.defs.lookup(x.ns + "/" + x.name) + return v, nil + case *future: + return nil, throwf("java.lang.UnsupportedOperationException", "proof serve runs futures once the eval is done, so it can't wait for one") + } + return nil, castError(args[0], "java.util.concurrent.Future") + }, + "Thread/sleep": func(e *evaluation, args []any) (any, error) { + if len(args) != 1 { + return nil, arity("Thread/sleep", len(args)) + } + ms, ok := args[0].(int64) + if !ok { + return nil, castError(args[0], "java.lang.Number") + } + if ms < 0 { + return nil, throwf("java.lang.IllegalArgumentException", "timeout value is negative") + } + t := time.NewTimer(time.Duration(min(ms, math.MaxInt64/int64(time.Millisecond))) * time.Millisecond) + defer t.Stop() + e.waiting() + select { + case <-t.C: + return nil, nil + case <-e.interrupts: + return nil, throwf("java.lang.InterruptedException", "sleep interrupted") + case <-e.ctx.Done(): + return nil, errStopping + } + }, + "apply": func(e *evaluation, args []any) (any, error) { + if len(args) < 2 { + return nil, arity("apply", len(args)) + } + var spread []any + switch last := args[len(args)-1].(type) { + case list: + spread = last + case vector: + spread = last + default: + return nil, castError(last, "clojure.lang.ISeq") + } + return e.apply(args[0], append(append([]any{}, args[1:len(args)-1]...), spread...)) + }, + "repeat": func(e *evaluation, args []any) (any, error) { + if len(args) != 2 { + return nil, arity("repeat", len(args)) + } + n, ok := args[0].(int64) + if !ok { + return nil, castError(args[0], "java.lang.Number") + } + // Plenty for a long value, while keeping one eval from eating + // all the memory. + if n > 1<<20 { + return nil, throwf("java.lang.OutOfMemoryError", "Java heap space") + } + xs := make(list, max(n, 0)) + for i := range xs { + xs[i] = args[1] + } + return xs, nil + }, + } +} + +// binder makes let and when-let, which bind names to values (only plain +// names, no destructuring). when-let gives up on a value that isn't +// truthy. +func binder(name string) builtin { + return func(e *evaluation, args []any) (any, error) { + var pairs vector + if len(args) > 0 { + pairs, _ = args[0].(vector) + } + if pairs == nil || len(pairs)%2 != 0 || (name == "when-let" && len(pairs) != 2) { + return nil, throwf("java.lang.IllegalArgumentException", "%s needs a vector of names and values", name) + } + saved := e.locals + defer func() { e.locals = saved }() + e.locals = maps.Clone(saved) + if e.locals == nil { + e.locals = map[string]any{} + } + for i := 0; i < len(pairs); i += 2 { + s, ok := pairs[i].(symbol) + if !ok { + return nil, throwf("java.lang.IllegalArgumentException", "proof serve only binds plain names") + } + v, err := e.eval(pairs[i+1]) + if err != nil || (name == "when-let" && !truthy(v)) { + return nil, err + } + e.locals[string(s)] = v + } + return e.do(args[1:]) + } +} + +func printer(readably, newline bool) builtin { + return func(e *evaluation, args []any) (any, error) { + parts := make([]string, len(args)) + for i, a := range args { + parts[i] = printed(a, readably) + } + text := strings.Join(parts, " ") + if newline { + text += "\n" + } + key := "out" + if e.toErr { + key = "err" + } + if text != "" { + e.output(key, text) + } + return nil, nil + } +} + +// arithmetic works on exact numbers like Clojure does, so (/ 1 2) is a +// ratio and (/ 1 0) throws. +func arithmetic(name string, op func(z, x, y *big.Rat) *big.Rat, identity int64) builtin { + return func(e *evaluation, args []any) (any, error) { + nums := make([]*big.Rat, len(args)) + for i, a := range args { + switch n := a.(type) { + case int64: + nums[i] = big.NewRat(n, 1) + case ratio: + nums[i] = big.NewRat(n.num, n.den) + default: + return nil, castError(a, "java.lang.Number") + } + } + if len(nums) == 0 { + if name == "-" || name == "/" { + return nil, arity(name, 0) + } + return identity, nil + } + acc, rest := new(big.Rat).Set(nums[0]), nums[1:] + if len(nums) == 1 && (name == "-" || name == "/") { + // (- x) is (- 0 x), and (/ x) is (/ 1 x). + acc, rest = big.NewRat(identity, 1), nums + } + for _, n := range rest { + if name == "/" && n.Sign() == 0 { + return nil, throwf("java.lang.ArithmeticException", "Divide by zero") + } + op(acc, acc, n) + } + if !acc.Num().IsInt64() || !acc.Denom().IsInt64() { + return nil, throwf("java.lang.ArithmeticException", "long overflow") + } + if acc.IsInt() { + return acc.Num().Int64(), nil + } + return ratio{acc.Num().Int64(), acc.Denom().Int64()}, nil + } +} + +func oneMore(e *evaluation, name, op string, args []any) (any, error) { + if len(args) != 1 { + return nil, arity(name, len(args)) + } + return builtins[op](e, []any{args[0], int64(1)}) +} + +func arity(name string, n int) error { + return throwf("clojure.lang.ArityException", "Wrong number of args (%d) passed to: clojure.core/%s", n, name) +} + +func castError(v any, to string) error { + return throwf("java.lang.ClassCastException", "class %s cannot be cast to class %s", className(v), to) +} + +func className(v any) string { + switch v.(type) { + case nil: + return "nil" + case bool: + return "java.lang.Boolean" + case int64: + return "java.lang.Long" + case ratio: + return "clojure.lang.Ratio" + case string: + return "java.lang.String" + case keyword: + return "clojure.lang.Keyword" + case symbol: + return "clojure.lang.Symbol" + case list: + return "clojure.lang.PersistentList" + case vector: + return "clojure.lang.PersistentVector" + case mapForm: + return "clojure.lang.PersistentArrayMap" + case namespace: + return "clojure.lang.Namespace" + case *variable: + return "clojure.lang.Var" + case *exception: + return "clojure.lang.ExceptionInfo" + } + return "java.lang.Object" +} diff --git a/internal/serve/eval_test.go b/internal/serve/eval_test.go new file mode 100644 index 0000000..e429add --- /dev/null +++ b/internal/serve/eval_test.go @@ -0,0 +1,39 @@ +package serve + +import ( + "context" + "testing" +) + +// The evaluator gets whatever code clients send, so nothing they send may +// crash it. +func FuzzEval(f *testing.F) { + for _, code := range []string{ + "(+ 1 2) (/ 1 0) (- 5) (/ 2)", `(str "a" nil [1 "b"] {:a 1}) (apply str (repeat 3 "ab"))`, + `(println "a" :k) (binding [*out* *err*] (prn 'x)) (flush)`, "(let [x 1] (when-let [y x] (if y x 2)))", + "(def x 1) #'x @(resolve 'x) (in-ns 'foo) (ns bar) *ns* *1 *e", `(throw (ex-info "x" {:a 1})) (ex-data *e)`, + "(read-line) (Thread/sleep 10) (future (println 1)) (require 'x)", "(let) (when) (1 2) '(1 \"a\") (quote)", + } { + f.Add(code) + } + interrupted := make(chan struct{}) + close(interrupted) + f.Fuzz(func(t *testing.T, code string) { + e := &evaluation{ + defs: newDefinitions(), b: newBindings(), ns: "user", ctx: context.Background(), + // Sleeps end right away, as interrupted. + interrupts: interrupted, + waiting: func() {}, + output: func(string, string) {}, + readLine: func(<-chan struct{}) (any, error) { return nil, nil }, + } + rd := &reader{src: code} + for { + form, err := rd.next() + if err != nil { + return + } + e.eval(form) + } + }) +} diff --git a/internal/serve/lang.go b/internal/serve/lang.go new file mode 100644 index 0000000..d17727b --- /dev/null +++ b/internal/serve/lang.go @@ -0,0 +1,297 @@ +package serve + +import ( + "fmt" + "io" + "strconv" + "strings" +) + +// The code proof serve evaluates is a sliver of Clojure: enough for the +// snippets in profiles/clojure.toml, and for client test suites to print, +// fail, read input and take their time. These are the forms it reads. +type ( + symbol string + keyword string + list []any + vector []any + // mapForm holds keys and values, alternating, in the order they were + // written. + mapForm []any +) + +// Everything else evaluation can produce. +type ( + ratio struct{ num, den int64 } + namespace string + variable struct{ ns, name string } + function struct{ name string } + future struct{} + // exception is what throw throws. data is ex-info's map. + exception struct { + class, msg string + data any + } +) + +// reader reads forms from code one at a time, the way nREPL does, so the +// forms before a broken one still get evaluated. +type reader struct { + src string + pos int +} + +// next returns the next form, or io.EOF once there are none left. +func (r *reader) next() (any, error) { + r.skip() + if r.pos >= len(r.src) { + return nil, io.EOF + } + return r.form() +} + +func (r *reader) skip() { + for r.pos < len(r.src) { + switch r.src[r.pos] { + case ' ', '\t', '\n', '\r', ',': + r.pos++ + case ';': + for r.pos < len(r.src) && r.src[r.pos] != '\n' { + r.pos++ + } + default: + return + } + } +} + +func (r *reader) form() (any, error) { + start := r.pos + switch c := r.src[r.pos]; c { + case '(': + items, err := r.until(')', start) + return list(items), err + case '[': + items, err := r.until(']', start) + return vector(items), err + case '{': + items, err := r.until('}', start) + if err == nil && len(items)%2 != 0 { + return nil, r.errorf("Map literal must contain an even number of forms") + } + return mapForm(items), err + case ')', ']', '}': + r.pos++ + return nil, r.errorf("Unmatched delimiter: %c", c) + case '"': + return r.str(start) + case '\'': + return r.wrapped("quote", 1, start) + case '@': + return r.wrapped("deref", 1, start) + case '#': + if strings.HasPrefix(r.src[r.pos:], "#'") { + return r.wrapped("var", 2, start) + } + r.pos++ + return nil, r.errorf("proof serve can't read %c", c) + case '`', '~', '^', '\\': + r.pos++ + return nil, r.errorf("proof serve can't read %c", c) + } + return r.token() +} + +// wrapped reads what a reader macro like ' applies to, after its n +// characters, and wraps it in a call to op. +func (r *reader) wrapped(op string, n, start int) (any, error) { + r.pos += n + r.skip() + if r.pos >= len(r.src) { + return nil, r.eof(start) + } + f, err := r.form() + return list{symbol(op), f}, err +} + +// until reads forms up to the closing delimiter. +func (r *reader) until(closing byte, start int) ([]any, error) { + r.pos++ + var items []any + for { + r.skip() + if r.pos >= len(r.src) { + return nil, r.eof(start) + } + if r.src[r.pos] == closing { + r.pos++ + return items, nil + } + f, err := r.form() + if err != nil { + return nil, err + } + items = append(items, f) + } +} + +func (r *reader) str(start int) (any, error) { + r.pos++ + var b strings.Builder + for r.pos < len(r.src) { + c := r.src[r.pos] + r.pos++ + switch { + case c == '"': + return b.String(), nil + case c != '\\': + b.WriteByte(c) + case r.pos >= len(r.src): + return nil, r.eof(start) + default: + e := r.src[r.pos] + r.pos++ + switch e { + case 'n': + b.WriteByte('\n') + case 't': + b.WriteByte('\t') + case 'r': + b.WriteByte('\r') + case '"', '\\': + b.WriteByte(e) + default: + return nil, r.errorf("Unsupported escape character: \\%c", e) + } + } + } + return nil, r.eofError("EOF while reading string") +} + +func (r *reader) token() (any, error) { + start := r.pos + for r.pos < len(r.src) && !strings.ContainsRune(" \t\n\r,;()[]{}\"", rune(r.src[r.pos])) { + r.pos++ + } + tok := r.src[start:r.pos] + switch tok { + case "nil": + return nil, nil + case "true": + return true, nil + case "false": + return false, nil + } + if k, ok := strings.CutPrefix(tok, ":"); ok { + if k == "" { + return nil, r.errorf("Invalid token: :") + } + return keyword(k), nil + } + if digits := strings.TrimLeft(tok, "+-"); digits != "" && digits[0] >= '0' && digits[0] <= '9' { + n, err := strconv.ParseInt(tok, 10, 64) + if err != nil { + return nil, r.errorf("Invalid number: %s", tok) + } + return n, nil + } + return symbol(tok), nil +} + +// eof reports code that ends in the middle of a form. +func (r *reader) eof(start int) error { + return r.eofError(fmt.Sprintf("EOF while reading, starting at line %d", 1+strings.Count(r.src[:start], "\n"))) +} + +// eofError reports the end of the code like nREPL does, at the line after +// it. +func (r *reader) eofError(msg string) error { + return readError(2+strings.Count(r.src, "\n"), 1, msg) +} + +func (r *reader) errorf(format string, args ...any) error { + before := r.src[:r.pos] + return readError(1+strings.Count(before, "\n"), 1+len(before)-(strings.LastIndexByte(before, '\n')+1), fmt.Sprintf(format, args...)) +} + +func readError(line, col int, msg string) error { + return &thrown{ex: &exception{class: "clojure.lang.ExceptionInfo", msg: msg}, phase: reading, line: line, col: col} +} + +// pr prints v the way pr-str does. +func pr(v any) string { return printed(v, true) } + +// printed prints v, with strings as they are unless readably, the way +// print does then. +func printed(v any, readably bool) string { + var b strings.Builder + write(&b, v, readably) + return b.String() +} + +func write(b *strings.Builder, v any, readably bool) { + items := func(open, close string, xs []any) { + b.WriteString(open) + for i, x := range xs { + if i > 0 { + b.WriteString(" ") + } + write(b, x, readably) + } + b.WriteString(close) + } + switch x := v.(type) { + case nil: + b.WriteString("nil") + case bool: + b.WriteString(strconv.FormatBool(x)) + case int64: + b.WriteString(strconv.FormatInt(x, 10)) + case ratio: + fmt.Fprintf(b, "%d/%d", x.num, x.den) + case string: + if readably { + b.WriteString(quote(x)) + } else { + b.WriteString(x) + } + case keyword: + b.WriteString(":" + string(x)) + case symbol: + b.WriteString(string(x)) + case list: + items("(", ")", x) + case vector: + items("[", "]", x) + case mapForm: + // Clojure separates the entries with commas: {:a 1, :b 2}. + b.WriteString("{") + for i := 0; i < len(x); i += 2 { + if i > 0 { + b.WriteString(", ") + } + write(b, x[i], readably) + b.WriteString(" ") + write(b, x[i+1], readably) + } + b.WriteString("}") + case namespace: + fmt.Fprintf(b, "#object[clojure.lang.Namespace 0x1 %q]", string(x)) + case *variable: + fmt.Fprintf(b, "#'%s/%s", x.ns, x.name) + case *function: + fmt.Fprintf(b, "#object[clojure.core$%s 0x1 \"clojure.core$%s@1\"]", x.name, x.name) + case *future: + b.WriteString("#object[clojure.core$future_call$reify__1 0x1 {:status :pending, :val nil}]") + case *exception: + fmt.Fprintf(b, "#error {\n :cause %s\n :data ", quote(x.msg)) + write(b, x.data, true) + fmt.Fprintf(b, "\n :via\n [{:type %s\n :message %s}]}", x.class, quote(x.msg)) + default: + panic(fmt.Sprintf("can't print %T", v)) + } +} + +var escapes = strings.NewReplacer(`\`, `\\`, `"`, `\"`, "\n", `\n`, "\t", `\t`, "\r", `\r`) + +func quote(s string) string { return `"` + escapes.Replace(s) + `"` } diff --git a/internal/serve/matrix_test.go b/internal/serve/matrix_test.go new file mode 100644 index 0000000..d3685a6 --- /dev/null +++ b/internal/serve/matrix_test.go @@ -0,0 +1,93 @@ +package serve + +import ( + "fmt" + "testing" + "time" + + "github.com/nrepl/proof/internal/check" + "github.com/nrepl/proof/internal/check/checktest" + "github.com/nrepl/proof/internal/checks" + "github.com/nrepl/proof/internal/profile" +) + +// serveFor starts a server with the given scenarios, stopped when the test +// ends. +func serveFor(t *testing.T, scenarios ...string) *Server { + t.Helper() + s, err := Listen("127.0.0.1:0", "0.1.0-dev", scenarios) + if err != nil { + t.Fatal(err) + } + go s.Serve() + t.Cleanup(func() { s.Stop(0) }) + return s +} + +// runChecks runs proof's server checks against proof serve, with the +// snippets of the profile for nREPL itself. +func runChecks(t *testing.T, scenarios ...string) map[string]check.Result { + t.Helper() + p, err := profile.Load("../../profiles/clojure.toml") + if err != nil { + t.Fatal(err) + } + p.Timeout = time.Second + env := &check.Env{Profile: p, Addr: serveFor(t, scenarios...).Addr(), Settle: 20 * time.Millisecond} + return checktest.ByID(check.Run(env, checks.All(), checks.WireRules())) +} + +func TestPassesEveryServerCheck(t *testing.T) { + checktest.Verdicts(t, runChecks(t), nil) +} + +// A scenario has to do what the servers it names do, so it gets the +// verdicts they get in the compatibility matrix. +func TestScenariosGetTheVerdictsOfTheirServers(t *testing.T) { + F, W, S := check.Failed, check.Warned, check.Skipped + noStdin := map[string]check.Verdict{"stdin.need-input": S, "stdin.roundtrip": S, "stdin.eof": S} + cases := []struct { + scenarios []string + want map[string]check.Verdict + }{ + {[]string{"split-output"}, nil}, + {[]string{"empty-messages"}, nil}, + {[]string{"last-value"}, map[string]check.Verdict{"eval.multiple-forms": F}}, + {[]string{"no-err"}, map[string]check.Verdict{"eval.stderr": F}}, + {[]string{"error-with-done"}, nil}, + {[]string{"no-op-echo"}, map[string]check.Verdict{"op.unknown-echo": W}}, + {[]string{"no-close-op"}, map[string]check.Verdict{"describe.required-ops": F}}, + {[]string{"no-interrupt"}, nil}, + {[]string{"no-stdin"}, noStdin}, + {[]string{"string-versions"}, nil}, + {[]string{"no-session-closed"}, map[string]check.Verdict{"session.close": F}}, + {[]string{"shared-state"}, map[string]check.Verdict{"session.isolated": F}}, + {[]string{"socket-sessions"}, map[string]check.Verdict{"session.across-connections": W}}, + {[]string{"any-session"}, map[string]check.Verdict{"session.unknown": F, "session.closed": F}}, + {[]string{"ns-fallback"}, map[string]check.Verdict{"eval.unknown-ns": F}}, + {[]string{"ns-error"}, map[string]check.Verdict{"eval.unknown-ns": F}}, + {[]string{"eof-error"}, map[string]check.Verdict{"stdin.eof": W}}, + {[]string{"unsorted-keys"}, map[string]check.Verdict{"wire.canonical": W}}, + {[]string{"byte-writes"}, nil}, + {[]string{"batched-writes"}, nil}, + // The columns of whole servers, except for eval.no-code, which no + // scenario covers as no client sends an eval without code. + {[]string{"no-op-echo", "socket-sessions", "shared-state", "no-session-closed", "any-session", "no-err", + "last-value", "ns-error", "no-stdin", "no-interrupt", "split-output"}, // Basilisp + map[string]check.Verdict{"op.unknown-echo": W, "session.across-connections": W, "session.isolated": F, + "session.close": F, "session.unknown": F, "session.closed": F, "eval.stderr": F, + "eval.multiple-forms": F, "eval.unknown-ns": F, "stdin.need-input": S, "stdin.roundtrip": S, "stdin.eof": S}}, + {[]string{"no-close-op", "no-op-echo", "socket-sessions", "shared-state", "no-session-closed", "any-session", + "last-value", "ns-fallback", "no-stdin", "no-interrupt", "unsorted-keys", "empty-messages", "error-with-done"}, // jank + map[string]check.Verdict{"describe.required-ops": F, "op.unknown-echo": W, "session.across-connections": W, + "session.isolated": F, "session.close": F, "session.unknown": F, "session.closed": F, + "eval.multiple-forms": F, "eval.unknown-ns": F, "stdin.need-input": S, "stdin.roundtrip": S, + "stdin.eof": S, "wire.canonical": W}}, + } + for _, c := range cases { + t.Run(fmt.Sprint(c.scenarios), func(t *testing.T) { + t.Parallel() + checktest.Verdicts(t, runChecks(t, c.scenarios...), c.want) + }) + } +} diff --git a/internal/serve/scenarios.go b/internal/serve/scenarios.go new file mode 100644 index 0000000..10dbde9 --- /dev/null +++ b/internal/serve/scenarios.go @@ -0,0 +1,110 @@ +package serve + +import ( + "fmt" + "slices" + "strings" +) + +// Scenario is something proof serve can do differently from nREPL on +// request: what some real server does, or what the network can do to the +// replies. Each one is taken from the compatibility matrix or from a +// server's own replies, never made up. +type Scenario struct { + Name string + Title string + // Who says who does it. + Who string +} + +// behavior is the set of scenarios a server runs with. +type behavior struct { + splitOutput, emptyMessages, lastValue, noErr, errorWithDone bool + noOpEcho, noCloseOp, noInterrupt, noStdin, stringVersions bool + noSessionClosed, sharedState, socketSessions, anySession bool + nsFallback, nsError, eofError bool + unsortedKeys, byteWrites, batchedWrites, hangUp bool +} + +// scenario is a Scenario and the behavior it sets. +type scenario struct { + Scenario + set func(*behavior) +} + +var catalog = []scenario{ + {Scenario{"split-output", "Output comes one character per message", + "nREPL splits long output, and Basilisp sends println's newline on its own"}, + func(b *behavior) { b.splitOutput = true }}, + {Scenario{"empty-messages", "Replies to eval include messages with nothing but id and session", "jank"}, + func(b *behavior) { b.emptyMessages = true }}, + {Scenario{"last-value", "Only the value of the last form is sent", "Basilisp, jank, dialtone"}, + func(b *behavior) { b.lastValue = true }}, + {Scenario{"no-err", "What the code prints to stderr never reaches the client", "Basilisp, dialtone, repartee"}, + func(b *behavior) { b.noErr = true }}, + {Scenario{"error-with-done", "eval-error comes in the same message as done", "jank"}, + func(b *behavior) { b.errorWithDone = true }}, + {Scenario{"no-op-echo", "Replies to unknown ops don't say which op it was", "ClojureCLR, Basilisp, jank, dialtone, repartee"}, + func(b *behavior) { b.noOpEcho = true }}, + {Scenario{"no-close-op", "describe doesn't list close, even though close works", "jank"}, + func(b *behavior) { b.noCloseOp = true }}, + {Scenario{"no-interrupt", "There's no interrupt op", "Basilisp, jank"}, + func(b *behavior) { b.noInterrupt = true }}, + {Scenario{"no-stdin", "There's no stdin op, and reading input gets an empty string right away", "Basilisp"}, + func(b *behavior) { b.noStdin = true }}, + {Scenario{"string-versions", "versions.proof is a plain string rather than a dict", "Babashka, for versions.babashka"}, + func(b *behavior) { b.stringVersions = true }}, + {Scenario{"no-session-closed", "close replies with done alone, without session-closed", "Basilisp, jank, dialtone, repartee"}, + func(b *behavior) { b.noSessionClosed = true }}, + {Scenario{"shared-state", "Sessions on the same connection share *1, *e and the current namespace", "ClojureCLR, Basilisp, jank"}, + func(b *behavior) { b.sharedState = true }}, + {Scenario{"socket-sessions", "A session only exists on the connection that cloned it", "ClojureCLR, Basilisp, jank"}, + func(b *behavior) { b.socketSessions = true }}, + {Scenario{"any-session", "Requests for sessions that don't exist run in a new session", "ClojureCLR, Basilisp, jank"}, + func(b *behavior) { b.anySession = true }}, + {Scenario{"ns-fallback", "An eval in a namespace that doesn't exist runs in the current one", "jank"}, + func(b *behavior) { b.nsFallback = true }}, + {Scenario{"ns-error", "An eval in a namespace that doesn't exist fails without namespace-not-found", "Basilisp"}, + func(b *behavior) { b.nsError = true }}, + {Scenario{"eof-error", "Reading past the end of input fails instead of returning nil", "nREPL 1.7.0"}, + func(b *behavior) { b.eofError = true }}, + {Scenario{"unsorted-keys", "The keys of reply dicts aren't sorted", "jank"}, + func(b *behavior) { b.unsortedKeys = true }}, + {Scenario{"byte-writes", "Replies are written a byte at a time", "any server, as TCP can deliver a message in pieces"}, + func(b *behavior) { b.byteWrites = true }}, + {Scenario{"batched-writes", "Replies are held back and written together until the eval waits or ends", "any server, as TCP can deliver several messages together"}, + func(b *behavior) { b.batchedWrites = true }}, + {Scenario{"hang-up", "The server closes the connection instead of answering an eval", "any server that crashes"}, + func(b *behavior) { b.hangUp = true }}, +} + +// Scenarios lists every scenario. +func Scenarios() []Scenario { + s := make([]Scenario, len(catalog)) + for i, c := range catalog { + s[i] = c.Scenario + } + return s +} + +// conflicts are scenarios that can't be combined. +var conflicts = [][2]string{{"ns-fallback", "ns-error"}, {"byte-writes", "batched-writes"}} + +func behave(names []string) (behavior, error) { + var b behavior + chosen := map[string]bool{} + for _, name := range names { + i := slices.IndexFunc(catalog, func(c scenario) bool { return c.Name == name }) + if i < 0 { + return b, fmt.Errorf("unknown scenario %q (proof list shows them all)", name) + } + catalog[i].set(&b) + chosen[name] = true + } + for _, pair := range conflicts { + if chosen[pair[0]] && chosen[pair[1]] { + return b, fmt.Errorf("scenarios %s don't go together", strings.Join(pair[:], " and ")) + } + } + return b, nil +} diff --git a/internal/serve/serve.go b/internal/serve/serve.go new file mode 100644 index 0000000..2cfaf4b --- /dev/null +++ b/internal/serve/serve.go @@ -0,0 +1,549 @@ +// Package serve is the server behind proof serve: a small nREPL server +// that behaves like the reference implementation, or, scenario by +// scenario, like the servers that differ from it. Client test suites run +// against it to make sure a client copes with all of them. It records +// every connection, so the requests can be graded like the proxy's. +package serve + +import ( + "bufio" + "bytes" + "context" + "errors" + "fmt" + "io" + "maps" + "net" + "slices" + "strconv" + "strings" + "sync" + "time" + + "github.com/nrepl/proof/bencode" + "github.com/nrepl/proof/internal/check" + "github.com/nrepl/proof/internal/clients" + "github.com/nrepl/proof/nrepl" +) + +// Server accepts clients and answers them. +type Server struct { + *clients.Listener + b behavior + version string + + defs *definitions + // ctx ends when proof stops, taking every eval still running with it. + ctx context.Context + stop context.CancelFunc + // work counts the goroutines running sessions, evals and futures. + work sync.WaitGroup + + mu sync.Mutex + sessions map[string]*session +} + +// Listen starts accepting clients on addr, with the given scenarios. Call +// Serve to handle them. version goes into describe's versions. +func Listen(addr, version string, scenarios []string) (*Server, error) { + b, err := behave(scenarios) + if err != nil { + return nil, err + } + ln, err := net.Listen("tcp", addr) + if err != nil { + return nil, err + } + return newServer(ln, version, b), nil +} + +func newServer(ln net.Listener, version string, b behavior) *Server { + ctx, stop := context.WithCancel(context.Background()) + return &Server{ + Listener: clients.New(ln), b: b, version: version, defs: newDefinitions(), + ctx: ctx, stop: stop, sessions: map[string]*session{}, + } +} + +// Serve accepts clients until Stop is called. +func (s *Server) Serve() { + s.Listener.Serve(func(cc *clients.Conn) { + s.Log("Connection %d opened", cc.N) + c := &conn{Conn: cc, srv: s, shared: newBindings()} + c.serve() + }) +} + +// Stop stops accepting clients and returns the transcript of every +// connection in which the client sent something. Clients that are still +// connected get up to grace to hang up on their own, and are then +// disconnected. +func (s *Server) Stop(grace time.Duration) []check.Traffic { + s.Listener.Stop(grace) + // No connection is left to start anything new. + s.stop() + s.work.Wait() + return s.Traffic() +} + +type conn struct { + *clients.Conn + srv *Server + // shared are the bindings of every session made on the connection, in + // shared-state. + shared *bindings + + // wmu keeps replies whole. + wmu sync.Mutex + // batch holds replies until the eval waits or ends, in batched-writes. + batch []byte +} + +func (c *conn) serve() { + dec := bencode.NewDecoder(bufio.NewReader(c.Client)) + for { + v, err := dec.Decode() + ev, isFrame := nrepl.DecodeEvent(nrepl.Sent, v, err) + if isFrame { + c.Record(ev) + } + switch { + case err == nil && ev.Msg != nil: + if c.srv.handle(c, ev.Msg) { + continue + } + case isFrame && !errors.Is(err, io.ErrUnexpectedEOF): + // There's no telling what the client meant, or where its next + // request starts. + c.hangUp() + case !errors.Is(err, net.ErrClosed): + // The end of the stream (maybe in the middle of a request), or a + // reset. + c.Record(nrepl.Event{Dir: nrepl.Sent, Time: time.Now(), Closed: true}) + c.srv.Log("Connection %d closed", c.N) + } + return + } +} + +// hangUp closes the connection from the server's side. +func (c *conn) hangUp() { + c.wmu.Lock() + defer c.wmu.Unlock() + c.Record(nrepl.Event{Dir: nrepl.Received, Time: time.Now(), Closed: true}) + c.Client.Close() + c.srv.Log("Connection %d closed by the server", c.N) +} + +// reply sends a reply to req, carrying its id and session. +func (c *conn) reply(req nrepl.Message, fields map[string]any) { + msg := map[string]any{} + for _, k := range []string{"id", "session"} { + if v, ok := req[k]; ok { + msg[k] = v + } + } + maps.Copy(msg, fields) + raw := c.srv.encode(msg) + c.wmu.Lock() + defer c.wmu.Unlock() + c.Record(nrepl.Event{Dir: nrepl.Received, Time: time.Now(), Msg: msg, Data: msg, Raw: raw}) + // A client that went away finds out soon enough, and its side of the + // connection records that. + switch { + case c.srv.b.byteWrites: + for i := range raw { + if _, err := c.Client.Write(raw[i : i+1]); err != nil { + return + } + } + case c.srv.b.batchedWrites: + c.batch = append(c.batch, raw...) + if _, ok := msg["status"]; ok { + c.flushLocked() + } + default: + c.Client.Write(raw) + } +} + +// flush writes the replies held back in batched-writes. +func (c *conn) flush() { + c.wmu.Lock() + defer c.wmu.Unlock() + c.flushLocked() +} + +func (c *conn) flushLocked() { + if len(c.batch) > 0 { + c.Client.Write(c.batch) + c.batch = c.batch[:0] + } +} + +// encode encodes a reply. Every value in it came from a request or from +// proof serve itself, so it can always be encoded. +func (s *Server) encode(msg map[string]any) []byte { + if !s.b.unsortedKeys { + raw, err := bencode.Marshal(msg) + if err != nil { + panic(err) + } + return raw + } + keys := slices.Sorted(maps.Keys(msg)) + slices.Reverse(keys) + var buf bytes.Buffer + buf.WriteByte('d') + for _, k := range keys { + for _, v := range []any{k, msg[k]} { + raw, err := bencode.Marshal(v) + if err != nil { + panic(err) + } + buf.Write(raw) + } + } + buf.WriteByte('e') + return buf.Bytes() +} + +var ( + done = []any{"done"} + unknownSession = map[string]any{"status": []any{"error", "unknown-session", "done"}} + ops = []string{"clone", "close", "describe", "eval", "interrupt", "stdin"} +) + +// offers reports whether the server supports an op. +func (s *Server) offers(op string) bool { + switch op { + case "interrupt": + return !s.b.noInterrupt + case "stdin": + return !s.b.noStdin + } + return slices.Contains(ops, op) +} + +// handle answers a request, and reports false if the server hung up. +func (s *Server) handle(c *conn, req nrepl.Message) bool { + sess, ok := s.session(c, req) + if !ok { + c.reply(req, unknownSession) + return true + } + if !req.Has("session") { + // Like nREPL, the replies say which session the request ran in, + // even when it's one made up for the request. + req = req.With(nrepl.Message{"session": sess.id}) + } + op := req.Str("op") + if !s.offers(op) { + fields := map[string]any{"status": []any{"error", "unknown-op", "done"}} + if v, ok := req["op"]; ok && !s.b.noOpEcho { + fields["op"] = v + } + c.reply(req, fields) + return true + } + switch op { + case "describe": + c.reply(req, s.describe()) + case "clone": + s.clone(c, req, sess) + case "close": + s.close(c, req, sess) + case "eval": + if s.b.hangUp { + c.hangUp() + return false + } + s.eval(c, req, sess) + case "stdin": + sess.give(req.Str("stdin")) + c.reply(req, map[string]any{"status": done}) + case "interrupt": + s.interrupt(c, req, sess) + } + return true +} + +func (s *Server) describe() map[string]any { + listed := map[string]any{} + for _, op := range ops { + if s.offers(op) && !(op == "close" && s.b.noCloseOp) { + listed[op] = map[string]any{} + } + } + var version any = s.version + if !s.b.stringVersions { + version = versionDict(s.version) + } + return map[string]any{ + "ops": listed, + "versions": map[string]any{"proof": version}, + "aux": map[string]any{"current-ns": "user"}, + "status": done, + } +} + +// versionDict describes a version like nREPL does, e.g. 0.1.0-dev as +// major 0, minor 1 and incremental 0. +func versionDict(v string) map[string]any { + d := map[string]any{"version-string": v} + numbers, _, _ := strings.Cut(v, "-") + for i, part := range strings.SplitN(numbers, ".", 3) { + if n, err := strconv.ParseInt(part, 10, 64); err == nil { + d[[]string{"major", "minor", "incremental"}[i]] = n + } + } + return d +} + +// session finds the session a request names. A request that doesn't name +// one gets a new ephemeral session, and so does one naming a session that +// doesn't exist in any-session. ok is false when there's no such session. +func (s *Server) session(c *conn, req nrepl.Message) (sess *session, ok bool) { + if !req.Has("session") { + return s.ephemeral(c), true + } + id, _ := req["session"].(string) + s.mu.Lock() + sess, ok = s.sessions[id] + s.mu.Unlock() + if ok && s.b.socketSessions && sess.conn != c { + ok = false + } + if !ok && s.b.anySession { + return s.ephemeral(c), true + } + return sess, ok +} + +func (s *Server) ephemeral(c *conn) *session { + b := newBindings() + if s.b.sharedState { + b = c.shared + } + return newSession(c, b, true) +} + +func (s *Server) clone(c *conn, req nrepl.Message, from *session) { + b := from.b.copy() + if s.b.sharedState { + b = c.shared + } + sess := newSession(c, b, false) + s.mu.Lock() + s.sessions[sess.id] = sess + s.mu.Unlock() + s.work.Add(1) + go func() { + defer s.work.Done() + s.run(sess) + }() + c.reply(req, map[string]any{"new-session": sess.id, "status": done}) +} + +func (s *Server) close(c *conn, req nrepl.Message, sess *session) { + s.mu.Lock() + if s.sessions[sess.id] == sess { + delete(s.sessions, sess.id) + } + s.mu.Unlock() + status := []any{"done", "session-closed"} + if s.b.noSessionClosed { + status = done + } + c.reply(req, map[string]any{"status": status}) + sess.close() +} + +func (s *Server) interrupt(c *conn, req nrepl.Message, sess *session) { + if sess.ephemeral { + c.reply(req, map[string]any{"status": []any{"error", "session-ephemeral", "done"}}) + return + } + id := "" + if req.Has("interrupt-id") { + id = fmt.Sprint(req["interrupt-id"]) + } + j, mismatch := sess.interrupt(id) + switch { + case mismatch: + c.reply(req, map[string]any{"status": []any{"error", "interrupt-id-mismatch", "done"}}) + return + case j == nil: + c.reply(req, map[string]any{"status": []any{"session-idle", "done"}}) + return + } + // Like nREPL, this tells the eval it's done before telling the + // interrupt, and the code finds out about it after that. + j.c.reply(j.req, map[string]any{"status": []any{"done", "interrupted"}}) + c.reply(req, map[string]any{"status": done}) + signal(j.interrupts) +} + +func (s *Server) eval(c *conn, req nrepl.Message, sess *session) { + reply := func(fields map[string]any) { c.reply(req, fields) } + switch req["code"].(type) { + case string: + case nil: + reply(map[string]any{"status": []any{"error", "no-code", "done"}}) + return + default: + reply(map[string]any{"status": []any{"error", "unknown-code-type", "done"}}) + return + } + if ns, ok := req["ns"].(string); ok && !s.defs.hasNS(ns) { + switch { + case s.b.nsError: + reply(s.fail(reply, &thrown{ex: &exception{class: "java.lang.Exception", msg: "No namespace: " + ns + " found"}}, "user")) + return + case !s.b.nsFallback: + reply(map[string]any{"status": []any{"error", "namespace-not-found", "done"}, "ns": ns}) + return + } + } + j := &job{c: c, req: req, interrupts: make(chan struct{}, 1)} + if !sess.ephemeral { + sess.queue(j) + return + } + s.work.Add(1) + go func() { + defer s.work.Done() + s.runEval(sess, j) + }() +} + +// fail reports an exception the way nREPL does, and returns the message +// that ends the eval. In error-with-done that message carries the error. +func (s *Server) fail(reply func(map[string]any), t *thrown, ns string) map[string]any { + text, ex, rootEx := errorReport(t, ns) + reply(map[string]any{"err": text}) + fields := map[string]any{"ex": ex, "root-ex": rootEx, "status": []any{"eval-error"}} + if s.b.errorWithDone { + fields["status"] = []any{"eval-error", "done"} + return fields + } + reply(fields) + return map[string]any{"status": done} +} + +// run carries out the session's evals one at a time, as nREPL does. +func (s *Server) run(sess *session) { + for { + j, ok := sess.next(s.ctx) + if !ok { + return + } + s.runEval(sess, j) + } +} + +func (s *Server) runEval(sess *session, j *job) { + c, req := j.c, j.req + reply := func(fields map[string]any) { c.reply(req, fields) } + if s.b.emptyMessages { + reply(nil) + } + ns := sess.b.currentNS() + if n, ok := req["ns"].(string); ok && s.defs.hasNS(n) { + ns = n + } + e := &evaluation{ + defs: s.defs, b: sess.b, ns: ns, ctx: s.ctx, interrupts: j.interrupts, waiting: c.flush, + output: func(key, text string) { + if key == "err" && s.b.noErr { + return + } + if !s.b.splitOutput { + reply(map[string]any{key: text}) + return + } + for _, r := range text { + reply(map[string]any{key: string(r)}) + } + }, + readLine: func(interrupts <-chan struct{}) (any, error) { + if s.b.noStdin { + // Basilisp's read-line gets an empty string right away. + return "", nil + } + line, ok, err := sess.readLine(s.ctx, interrupts, func() { reply(map[string]any{"status": []any{"need-input"}}) }) + switch { + case err != nil: + return nil, err + case ok: + return line, nil + case s.b.eofError: + // What nREPL 1.7.0 says. + return nil, throwf("java.lang.ClassCastException", "class java.lang.Long cannot be cast to class java.lang.Character") + } + return nil, nil + }, + } + final := map[string]any{"status": done} + // held is the value last-value sends at the end. + var held map[string]any + rd := &reader{src: req.Str("code")} + for { + form, err := rd.next() + if err == io.EOF { + break + } + var v any + if err == nil { + v, err = e.eval(form) + } + if errors.Is(err, errStopping) { + return + } + if err != nil { + t := err.(*thrown) + sess.b.failed(t.ex) + final, held = s.fail(reply, t, e.ns), nil + // nREPL goes on with the next form, unless there's no telling + // where it starts. Basilisp and jank give up on the rest. + if t.phase == reading || s.b.lastValue || s.b.errorWithDone { + break + } + continue + } + sess.b.push(v) + msg := map[string]any{"value": pr(v), "ns": e.ns} + if s.b.lastValue { + held = msg + } else { + reply(msg) + } + } + if held != nil { + reply(held) + } + // A namespace given with the request is only for this eval. + if !req.Has("ns") { + sess.b.setNS(e.ns) + } + if sess.finish(j) { + reply(final) + } + // Futures print after the eval is done, so their output is late, as it + // is on nREPL. + s.runFutures(c, e.futures) +} + +func (s *Server) runFutures(c *conn, futures []pending) { + for _, f := range futures { + s.work.Add(1) + go func() { + defer s.work.Done() + // Interrupting the eval doesn't stop its futures. + f.e.interrupts = nil + f.e.do(f.body) + c.flush() + s.runFutures(c, f.e.futures) + }() + } +} diff --git a/internal/serve/serve_test.go b/internal/serve/serve_test.go new file mode 100644 index 0000000..4e4a99e --- /dev/null +++ b/internal/serve/serve_test.go @@ -0,0 +1,590 @@ +package serve + +import ( + "io" + "maps" + "net" + "reflect" + "slices" + "strings" + "sync" + "testing" + "time" + + "github.com/nrepl/proof/internal/check" + "github.com/nrepl/proof/internal/check/checktest" + "github.com/nrepl/proof/internal/checks" + "github.com/nrepl/proof/nrepl" +) + +const timeout = 5 * time.Second + +func connect(t *testing.T, s *Server) *nrepl.Conn { + t.Helper() + c, err := nrepl.Dial(s.Addr(), timeout) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { c.Close() }) + return c +} + +func request(t *testing.T, c *nrepl.Conn, m nrepl.Message) nrepl.Response { + t.Helper() + r, err := c.Request(m, timeout) + if err != nil { + t.Fatalf("%v: %v", m, err) + } + return r +} + +func clone(t *testing.T, c *nrepl.Conn) string { + t.Helper() + return request(t, c, nrepl.Message{"op": "clone"}).Str("new-session") +} + +func eval(t *testing.T, c *nrepl.Conn, session, code string) nrepl.Response { + t.Helper() + return request(t, c, nrepl.Message{"op": "eval", "code": code, "session": session}) +} + +func TestEvaluatesLikeClojure(t *testing.T) { + s := serveFor(t) + cases := []struct { + code string + values []string + out string + err string + }{ + {code: "(+ 1 2)", values: []string{"3"}}, + {code: "1 2", values: []string{"1", "2"}}, + {code: `:k "s\n" nil true [1 "a"] {:a 1 :b "c"} '(1 x)`, + values: []string{":k", `"s\n"`, "nil", "true", `[1 "a"]`, `{:a 1, :b "c"}`, "(1 x)"}}, + {code: "(- 5) (/ 1 2) (/ 4 2) (/ 2) (* 2 3 4) (inc 1) (dec 1)", + values: []string{"-5", "1/2", "2", "1/2", "24", "2", "0"}}, + {code: `(str "a" nil 1 :k) (apply str (repeat 3 "ab"))`, values: []string{`"a1:k"`, `"ababab"`}}, + {code: "(if nil 1 2) (do 1 2)", values: []string{"2", "2"}}, + {code: `(print "proof")`, values: []string{"nil"}, out: "proof"}, + {code: `(println "a" 1 :k nil ["x"]) (prn "a" {:a "b"})`, values: []string{"nil", "nil"}, + out: "a 1 :k nil [x]\n\"a\" {:a \"b\"}\n"}, + {code: `(binding [*out* *err*] (print "proof"))`, values: []string{"nil"}, err: "proof"}, + {code: "(def x 1) x user/x", values: []string{"#'user/x", "1", "1"}}, + {code: "(let [x 1 y (+ x 1)] y) (when 1 2) (when nil 2) (when-let [x nil] 1) (when-let [x 3] x)", + values: []string{"2", "2", "nil", "nil", "3"}}, + {code: "(require 'clojure.stacktrace) (def y 1) (resolve 'y) (resolve 'nope) @(resolve 'y) @#'y", + values: []string{"nil", "#'user/y", "#'user/y", "nil", "1", "1"}}, + {code: `(ex-data (ex-info "x" {:a 1}))`, values: []string{"{:a 1}"}}, + {code: "(clojure.core// 4 2) (resolve '/) (resolve 'clojure.core/str)", values: []string{"2", "#'clojure.core//", "#'clojure.core/str"}}, + {code: "#'nope", err: "Syntax error compiling var at (REPL:0:0).\nUnable to resolve var: nope in this context\n"}, + {code: "@(future 1)", err: "Execution error (UnsupportedOperationException) at user/eval1 (REPL:1).\n" + + "proof serve runs futures once the eval is done, so it can't wait for one\n"}, + {code: "(/ 1 0) 42", values: []string{"42"}, err: "Execution error (ArithmeticException) at user/eval1 (REPL:1).\nDivide by zero\n"}, + {code: `(str [1 "a"] {:a "b"} '(1 "c") *ns* #'clojure.core/str)`, + values: []string{`"[1 \"a\"]{:a \"b\"}(1 \"c\")user#'clojure.core/str"`}}, + {code: "(+ 9223372036854775807 1)", err: "Execution error (ArithmeticException) at user/eval1 (REPL:1).\nlong overflow\n"}, + {code: "(Thread/sleep -1)", err: "Execution error (IllegalArgumentException) at user/eval1 (REPL:1).\ntimeout value is negative\n"}, + {code: "(let)", err: "Execution error (IllegalArgumentException) at user/eval1 (REPL:1).\nlet needs a vector of names and values\n"}, + {code: "(let [z 1] z) z", values: []string{"1"}, err: "Syntax error compiling at (REPL:0:0).\nUnable to resolve symbol: z in this context\n"}, + {code: "(in-ns 'clojure.core) (def zz 5) @#'zz", values: []string{`#object[clojure.lang.Namespace 0x1 "clojure.core"]`, + "#'clojure.core/zz", "5"}}, + {code: `1 "abc`, values: []string{"1"}, err: "Syntax error reading source at (REPL:2:1).\nEOF while reading string\n"}, + {code: "(in-ns 'foo) *ns*", values: []string{`#object[clojure.lang.Namespace 0x1 "foo"]`, + `#object[clojure.lang.Namespace 0x1 "foo"]`}}, + {code: "(/ 1 0)", err: "Execution error (ArithmeticException) at user/eval1 (REPL:1).\nDivide by zero\n"}, + {code: `(throw (ex-info "proof" {}))`, err: "Execution error (ExceptionInfo) at user/eval1 (REPL:1).\nproof\n"}, + {code: `(+ 1 "a")`, err: "Execution error (ClassCastException) at user/eval1 (REPL:1).\n" + + "class java.lang.String cannot be cast to class java.lang.Number\n"}, + {code: "(1 2)", err: "Execution error (ClassCastException) at user/eval1 (REPL:1).\n" + + "class java.lang.Long cannot be cast to class clojure.lang.IFn\n"}, + {code: "foo", err: "Syntax error compiling at (REPL:0:0).\nUnable to resolve symbol: foo in this context\n"}, + {code: "(+ 1", err: "Syntax error reading source at (REPL:2:1).\nEOF while reading, starting at line 1\n"}, + {code: "1 )", values: []string{"1"}, err: "Syntax error reading source at (REPL:1:4).\nUnmatched delimiter: )\n"}, + {code: "#{1}", err: "Syntax error reading source at (REPL:1:2).\nproof serve can't read #\n"}, + } + c := connect(t, s) + for _, tc := range cases { + r := eval(t, c, clone(t, c), tc.code) + if got := r.Values(); !slices.Equal(got, tc.values) { + t.Errorf("%s: values %q, want %q", tc.code, got, tc.values) + } + if r.Out() != tc.out || r.Err() != tc.err { + t.Errorf("%s: out %q and err %q, want %q and %q", tc.code, r.Out(), r.Err(), tc.out, tc.err) + } + if failed := r.HasStatus("eval-error"); failed != (tc.err != "" && !strings.HasPrefix(tc.code, "(binding")) { + t.Errorf("%s: status %v", tc.code, r.Status()) + } + } +} + +func TestSessionsKeepTheirBindings(t *testing.T) { + c := connect(t, serveFor(t)) + first, second := clone(t, c), clone(t, c) + eval(t, c, first, ":marker (in-ns 'foo)") + if v := eval(t, c, first, "*2 *ns*").Values(); !slices.Equal(v, []string{":marker", `#object[clojure.lang.Namespace 0x1 "foo"]`}) { + t.Errorf("first session: %q", v) + } + if v := eval(t, c, second, "*1 *ns*").Values(); !slices.Equal(v, []string{"nil", `#object[clojure.lang.Namespace 0x1 "user"]`}) { + t.Errorf("second session: %q", v) + } + r := eval(t, c, first, "(throw (ex-info \"proof\" {:a 1}))") + if !r.HasStatus("eval-error") { + t.Fatal(r.Status()) + } + if v := eval(t, c, first, "*e").Values(); len(v) != 1 || !strings.Contains(v[0], ":cause \"proof\"\n :data {:a 1}") { + t.Errorf("*e: %q", v) + } + eval(t, c, first, "(+ 1") + if v := eval(t, c, first, "*e").Values(); len(v) != 1 || !strings.Contains(v[0], `:cause "EOF while reading`) { + t.Errorf("a read error didn't set *e: %q", v) + } + // A namespace sent with an eval is only for that eval. + in := func(ns, code string) []string { + return request(t, c, nrepl.Message{"op": "eval", "code": code, "session": second, "ns": ns}).Values() + } + in("user", "(in-ns 'qqq)") + if v := in("qqq", "(str *ns*)"); !slices.Equal(v, []string{`"qqq"`}) { + t.Errorf("eval in qqq: %q", v) + } + if v := eval(t, c, second, "(str *ns*)").Values(); !slices.Equal(v, []string{`"user"`}) { + t.Errorf("after evals with a namespace: %q", v) + } + // A clone starts with the bindings of the session it's cloned from. + cloned := request(t, c, nrepl.Message{"op": "clone", "session": first}).Str("new-session") + if v := eval(t, c, cloned, "*ns*").Values(); !slices.Equal(v, []string{`#object[clojure.lang.Namespace 0x1 "foo"]`}) { + t.Errorf("clone: %q", v) + } +} + +func TestFutureOutputComesAfterDone(t *testing.T) { + c := connect(t, serveFor(t)) + session := clone(t, c) + r := eval(t, c, session, `(let [s "late"] (future (Thread/sleep 10) (println s)))`) + if !strings.Contains(r.Values()[0], "future_call") { + t.Errorf("value %q", r.Values()) + } + late, err := c.WaitFor(func(m nrepl.Message) bool { return m.Str("out") == "late\n" }, timeout) + if err != nil { + t.Fatal(err) + } + id := r.Messages[0].Str("id") + if late.Str("id") != id || late.Str("session") != session { + t.Errorf("late output %v doesn't belong to the eval", late) + } + if got := kinds(c, id); !slices.Equal(got, []string{"value", "done", "out"}) { + t.Errorf("the eval got %v, want the output after done", got) + } +} + +func TestInput(t *testing.T) { + cases := []struct { + scenario string + inputs []string + values []string + err string + }{ + {inputs: []string{"proof\n", ""}, values: []string{`"proof"`, "nil"}}, + {inputs: []string{"pro", "of\nmore\n"}, values: []string{`"proof"`, `"more"`}}, + {inputs: []string{"partial", "", ""}, values: []string{`"partial"`, "nil"}}, + {inputs: []string{"", ""}, values: []string{"nil", "nil"}}, + {scenario: "eof-error", inputs: []string{"", ""}, + err: strings.Repeat("Execution error (ClassCastException) at user/eval1 (REPL:1).\n"+ + "class java.lang.Long cannot be cast to class java.lang.Character\n", 2)}, + {scenario: "no-stdin", values: []string{`""`, `""`}}, + } + for _, tc := range cases { + t.Run(strings.Join(append([]string{tc.scenario}, tc.inputs...), "|"), func(t *testing.T) { + var scenarios []string + if tc.scenario != "" { + scenarios = append(scenarios, tc.scenario) + } + c := connect(t, serveFor(t, scenarios...)) + session := clone(t, c) + id, err := c.Send(nrepl.Message{"op": "eval", "code": "(read-line) (read-line)", "session": session}) + if err != nil { + t.Fatal(err) + } + for _, input := range tc.inputs { + if _, err := c.WaitFor(func(m nrepl.Message) bool { return m.HasStatus("need-input") }, timeout); err != nil { + t.Fatal(err) + } + request(t, c, nrepl.Message{"op": "stdin", "stdin": input, "session": session}) + } + r, err := c.Collect(id, timeout) + if err != nil { + t.Fatal(err) + } + if !slices.Equal(r.Values(), tc.values) || r.Err() != tc.err { + t.Errorf("values %q and err %q, want %q and %q", r.Values(), r.Err(), tc.values, tc.err) + } + }) + } +} + +// received returns what the server sent for a request so far, done or not. +func received(c *nrepl.Conn, id string) nrepl.Response { + var r nrepl.Response + for _, ev := range c.Transcript() { + if ev.Dir == nrepl.Received && ev.Msg.Str("id") == id { + r.Messages = append(r.Messages, ev.Msg) + } + } + return r +} + +// kinds says what each message the server sent for a request so far was, +// in order. +func kinds(c *nrepl.Conn, id string) []string { + var ks []string + for _, m := range received(c, id).Messages { + for _, k := range []string{"done", "need-input", "eval-error"} { + if m.HasStatus(k) { + ks = append(ks, k) + } + } + for _, k := range []string{"out", "err", "value"} { + if m.Has(k) { + ks = append(ks, k) + } + } + } + return ks +} + +func TestInterrupt(t *testing.T) { + c := connect(t, serveFor(t)) + session := clone(t, c) + if r := request(t, c, nrepl.Message{"op": "interrupt", "session": session}); !r.HasStatus("session-idle", "done") { + t.Errorf("idle session: %v", r.Status()) + } + if r := request(t, c, nrepl.Message{"op": "interrupt"}); !r.HasStatus("error", "session-ephemeral", "done") { + t.Errorf("no session: %v", r.Status()) + } + // What each of them sends once it's running. + running := map[string]string{ + `(do (print "sleeping") (Thread/sleep 60000)) :after`: "out", + "(read-line) :after": "need-input", + } + for code, started := range running { + id, err := c.Send(nrepl.Message{"op": "eval", "code": code, "session": session}) + if err != nil { + t.Fatal(err) + } + if _, err := c.WaitFor(func(m nrepl.Message) bool { return m.Str("id") == id }, timeout); err != nil { + t.Fatal(err) + } + r := request(t, c, nrepl.Message{"op": "interrupt", "session": session, "interrupt-id": "something else"}) + if !r.HasStatus("error", "interrupt-id-mismatch", "done") { + t.Errorf("%s: interrupting some other eval: %v", code, r.Status()) + } + interrupt, err := c.Send(nrepl.Message{"op": "interrupt", "session": session, "interrupt-id": id}) + if err != nil { + t.Fatal(err) + } + // Like on nREPL, the eval is done before the interrupt, and the rest + // of its code runs after that, without another done. + if _, err := c.WaitFor(func(m nrepl.Message) bool { return m.Str("id") == id && m.Has("value") }, timeout); err != nil { + t.Fatal(err) + } + c.Settle(50*time.Millisecond, time.Second) + if got, want := kinds(c, id), []string{started, "done", "err", "eval-error", "value"}; !slices.Equal(got, want) { + t.Errorf("%s: the eval got %v, want %v", code, got, want) + } + var order []string + for _, ev := range c.Transcript() { + if ev.Msg.HasStatus("done") && (ev.Msg.Str("id") == id || ev.Msg.Str("id") == interrupt) { + order = append(order, ev.Msg.Str("id")) + } + } + if !slices.Equal(order, []string{id, interrupt}) { + t.Errorf("%s: dones came in the order %v, want the eval's first", code, order) + } + if all := received(c, id); !all.HasStatus("done", "interrupted") || !strings.Contains(all.Err(), "(InterruptedException)") { + t.Errorf("%s: interrupted eval: %v %q", code, all.Status(), all.Err()) + } + } + // The session goes on. + if v := eval(t, c, session, "(+ 1 2)").Values(); !slices.Equal(v, []string{"3"}) { + t.Errorf("after the interrupts: %q", v) + } +} + +func TestClosingASessionInterruptsItsEval(t *testing.T) { + c := connect(t, serveFor(t)) + session := clone(t, c) + id, err := c.Send(nrepl.Message{"op": "eval", "code": `(do (print "sleeping") (Thread/sleep 60000)) :after`, "session": session}) + if err != nil { + t.Fatal(err) + } + if _, err := c.WaitFor(func(m nrepl.Message) bool { return m.Str("out") == "sleeping" }, timeout); err != nil { + t.Fatal(err) + } + if r := request(t, c, nrepl.Message{"op": "close", "session": session}); !r.HasStatus("session-closed") { + t.Errorf("close: %v", r.Status()) + } + r, err := c.Collect(id, timeout) + if err != nil { + t.Fatal(err) + } + if !strings.Contains(r.Err(), "(InterruptedException)") || !slices.Equal(r.Values(), []string{":after"}) || r.HasStatus("interrupted") { + t.Errorf("the eval got %v", r.Messages) + } +} + +// The scenarios checked here don't change any verdict of the server +// checks, so the matrix test can't tell whether they work. +func TestScenarioShapes(t *testing.T) { + cases := []struct { + scenario string + code string + want func(r nrepl.Response) bool + }{ + {"split-output", `(print "aä€")`, func(r nrepl.Response) bool { + var chunks []string + for _, m := range r.Messages { + if m.Has("out") { + chunks = append(chunks, m.Str("out")) + } + } + return slices.Equal(chunks, []string{"a", "ä", "€"}) + }}, + {"empty-messages", "1", func(r nrepl.Response) bool { + return slices.ContainsFunc(r.Messages, func(m nrepl.Message) bool { return len(m) == 2 && m.Has("id") && m.Has("session") }) + }}, + {"error-with-done", "(/ 1 0) 2", func(r nrepl.Response) bool { + last := r.Messages[len(r.Messages)-1] + return len(r.Messages) == 2 && last.HasStatus("eval-error", "done") && last.Has("ex") + }}, + {"no-err", `(binding [*out* *err*] (print "x")) (/ 1 0)`, func(r nrepl.Response) bool { + // Only the error report is left. + return strings.HasPrefix(r.Err(), "Execution error") + }}, + {"last-value", "1 2 3", func(r nrepl.Response) bool { return slices.Equal(r.Values(), []string{"3"}) }}, + {"last-value", "1 (/ 1 0) 3", func(r nrepl.Response) bool { return len(r.Values()) == 0 && r.HasStatus("eval-error") }}, + } + for _, tc := range cases { + t.Run(tc.scenario, func(t *testing.T) { + c := connect(t, serveFor(t, tc.scenario)) + if r := eval(t, c, clone(t, c), tc.code); !tc.want(r) { + t.Errorf("unexpected replies: %v", r.Messages) + } + }) + } +} + +func TestDescribe(t *testing.T) { + cases := []struct { + scenarios []string + ops []string + version any + }{ + {nil, []string{"clone", "close", "describe", "eval", "interrupt", "stdin"}, + map[string]any{"major": int64(0), "minor": int64(1), "incremental": int64(0), "version-string": "0.1.0-dev"}}, + {[]string{"no-close-op", "no-interrupt", "no-stdin", "string-versions"}, []string{"clone", "describe", "eval"}, "0.1.0-dev"}, + } + for _, tc := range cases { + c := connect(t, serveFor(t, tc.scenarios...)) + r := request(t, c, nrepl.Message{"op": "describe"}) + ops, _ := r.Get("ops").(map[string]any) + if got := slices.Sorted(maps.Keys(ops)); !slices.Equal(got, tc.ops) { + t.Errorf("%v: ops %v, want %v", tc.scenarios, got, tc.ops) + } + versions, _ := r.Get("versions").(map[string]any) + if !reflect.DeepEqual(versions["proof"], tc.version) { + t.Errorf("%v: versions %v", tc.scenarios, versions) + } + } +} + +func TestOpsTheScenariosTakeAway(t *testing.T) { + c := connect(t, serveFor(t, "no-interrupt", "no-stdin")) + for _, op := range []string{"interrupt", "stdin"} { + if r := request(t, c, nrepl.Message{"op": op, "session": clone(t, c)}); !r.HasStatus("unknown-op") { + t.Errorf("%s: %v", op, r.Status()) + } + } +} + +// pipeListener hands out the server's end of a pipe, where every write +// arrives on its own. +type pipeListener struct { + conns chan net.Conn + closed chan struct{} + once sync.Once +} + +func (l *pipeListener) Accept() (net.Conn, error) { + select { + case c := <-l.conns: + return c, nil + case <-l.closed: + return nil, net.ErrClosed + } +} + +func (l *pipeListener) Close() error { + l.once.Do(func() { close(l.closed) }) + return nil +} + +func (l *pipeListener) Addr() net.Addr { return &net.TCPAddr{} } + +// writes connects a client to a server over a pipe, and returns what each +// write of the replies to an eval carried. +func writes(t *testing.T, scenario string) []string { + t.Helper() + b, err := behave([]string{scenario}) + if err != nil { + t.Fatal(err) + } + client, server := net.Pipe() + ln := &pipeListener{conns: make(chan net.Conn, 1), closed: make(chan struct{})} + ln.conns <- server + s := newServer(ln, "0.1.0-dev", b) + go s.Serve() + t.Cleanup(func() { + client.Close() + s.Stop(0) + }) + go client.Write([]byte("d4:code11:(print 1) 22:id1:12:op4:evale")) + var got []string + buf := make([]byte, 1024) + client.SetReadDeadline(time.Now().Add(timeout)) + for !strings.Contains(strings.Join(got, ""), "4:done") { + n, err := client.Read(buf) + if err != nil { + t.Fatalf("after %q: %v", got, err) + } + got = append(got, string(buf[:n])) + } + return got +} + +func TestByteWrites(t *testing.T) { + for _, w := range writes(t, "byte-writes") { + if len(w) != 1 { + t.Fatalf("a write of %q", w) + } + } +} + +func TestBatchedWrites(t *testing.T) { + // One write for everything up to the done. + if w := writes(t, "batched-writes"); len(w) != 1 || strings.Count(w[0], "2:id") != 4 { + t.Errorf("writes %q", w) + } +} + +func TestHangUp(t *testing.T) { + s := serveFor(t, "hang-up") + c := connect(t, s) + request(t, c, nrepl.Message{"op": "describe"}) + if _, err := c.Request(nrepl.Message{"op": "eval", "code": "1"}, timeout); err == nil { + t.Fatal("the eval got its done") + } + traffic := s.Stop(timeout) + if events := traffic[0].Events; !events[len(events)-1].Closed || events[len(events)-1].Dir != nrepl.Received { + t.Errorf("the server's hanging up isn't recorded last: %v", events[len(events)-1]) + } +} + +func TestStopGradesWhatClientsSent(t *testing.T) { + s := serveFor(t) + c := connect(t, s) + session := clone(t, c) + // Still sleeping when proof stops, which doesn't hold Stop up. + c.Send(nrepl.Message{"op": "eval", "code": "(Thread/sleep 60000)", "session": session}) + c.Close() + + start := time.Now() + traffic := s.Stop(timeout) + if elapsed := time.Since(start); elapsed > time.Second { + t.Errorf("Stop took %s", elapsed) + } + if len(traffic) != 1 { + t.Fatalf("traffic: %v", traffic) + } + if v := checktest.ByID(check.Grade(checks.ClientRules(), traffic))["client.close"].Verdict; v != check.Warned { + t.Errorf("the session was never closed, but client.close got %s", v) + } +} + +func TestBrokenFramesEndTheConnection(t *testing.T) { + s := serveFor(t) + nc, err := net.Dial("tcp", s.Addr()) + if err != nil { + t.Fatal(err) + } + defer nc.Close() + nc.Write([]byte("d2:op8:describeXe")) + nc.SetReadDeadline(time.Now().Add(timeout)) + if _, err := io.ReadAll(nc); err != nil { + t.Errorf("the server didn't hang up: %v", err) + } +} + +func TestScenarioNames(t *testing.T) { + for _, names := range [][]string{{"no-such-thing"}, {"ns-fallback", "ns-error"}, {"byte-writes", "batched-writes"}} { + if _, err := Listen("127.0.0.1:0", "0", names); err == nil { + t.Errorf("%v: no error", names) + } + } + for _, sc := range Scenarios() { + if _, err := behave([]string{sc.Name}); err != nil { + t.Error(err) + } + } +} + +func TestRepliesSayWhichSessionTheyRanIn(t *testing.T) { + c := connect(t, serveFor(t)) + sessions := map[string]bool{} + for _, req := range []nrepl.Message{{"op": "describe"}, {"op": "eval", "code": "(print 1) 2"}, {"op": "proof/nope"}} { + r := request(t, c, req) + for _, m := range r.Messages { + if m.Str("session") != r.Messages[0].Str("session") || m.Str("session") == "" { + t.Errorf("%v: replies in different sessions %v", req, r.Messages) + } + } + sessions[r.Messages[0].Str("session")] = true + } + if len(sessions) != 3 { + t.Errorf("each request should get a session of its own, got %v", sessions) + } +} + +func TestCodeThatIsntAString(t *testing.T) { + c := connect(t, serveFor(t)) + if r := request(t, c, nrepl.Message{"op": "eval", "code": int64(1)}); !r.HasStatus("error", "unknown-code-type", "done") { + t.Errorf("status %v", r.Status()) + } +} + +func TestBatchedWritesDontHoldOutputBack(t *testing.T) { + c := connect(t, serveFor(t, "batched-writes")) + c.Send(nrepl.Message{"op": "eval", "code": `(do (print "sleeping") (Thread/sleep 60000))`, "session": clone(t, c)}) + if _, err := c.WaitFor(func(m nrepl.Message) bool { return m.Str("out") == "sleeping" }, timeout); err != nil { + t.Errorf("output before a sleep: %v", err) + } +} + +func TestHangingUpInTheMiddleOfARequest(t *testing.T) { + s := serveFor(t) + nc, err := net.Dial("tcp", s.Addr()) + if err != nil { + t.Fatal(err) + } + nc.Write([]byte("d2:id1:12:op5:clonee")) + nc.SetReadDeadline(time.Now().Add(timeout)) + if _, err := nc.Read(make([]byte, 1024)); err != nil { + t.Fatal(err) + } + nc.Write([]byte("d2:op4:eval")) + nc.Close() + traffic := s.Stop(timeout) + if events := traffic[0].Events; !events[len(events)-1].Closed || events[len(events)-1].Dir != nrepl.Sent { + t.Errorf("the client's hanging up isn't recorded last: %v", events[len(events)-1]) + } + if v := checktest.ByID(check.Grade(checks.ClientRules(), traffic))["client.close"].Verdict; v != check.Warned { + t.Errorf("the session was never closed, but client.close got %s", v) + } +} diff --git a/internal/serve/session.go b/internal/serve/session.go new file mode 100644 index 0000000..cd06189 --- /dev/null +++ b/internal/serve/session.go @@ -0,0 +1,184 @@ +package serve + +import ( + "context" + "crypto/rand" + "fmt" + "strings" + "sync" + + "github.com/nrepl/proof/nrepl" +) + +// job is an eval request, waiting for its turn or running. +type job struct { + c *conn + req nrepl.Message + // interrupts works like a thread's interrupt flag: the next sleep or + // read of input gets an InterruptedException. + interrupts chan struct{} + // interrupted means interrupt has sent the eval's done already. + interrupted bool +} + +type session struct { + id string + // conn is the connection that created the session. + conn *conn + b *bindings + // ephemeral sessions only last for one request. + ephemeral bool + + mu sync.Mutex + jobs []*job + running *job + closed bool + // input is what stdin sent that hasn't been read yet, and eof means + // an empty stdin came after it. + input string + eof bool + // wake says there's a job or the session was closed, and inputs that + // input arrived. + wake, inputs chan struct{} +} + +func newSession(c *conn, b *bindings, ephemeral bool) *session { + return &session{ + id: newID(), conn: c, b: b, ephemeral: ephemeral, + wake: make(chan struct{}, 1), inputs: make(chan struct{}, 1), + } +} + +// newID makes a session id like nREPL's, a random UUID. +func newID() string { + b := make([]byte, 16) + rand.Read(b) + b[6] = b[6]&0x0f | 0x40 + b[8] = b[8]&0x3f | 0x80 + return fmt.Sprintf("%x-%x-%x-%x-%x", b[0:4], b[4:6], b[6:8], b[8:10], b[10:]) +} + +func signal(ch chan struct{}) { + select { + case ch <- struct{}{}: + default: + } +} + +func (sess *session) queue(j *job) { + sess.mu.Lock() + sess.jobs = append(sess.jobs, j) + sess.mu.Unlock() + signal(sess.wake) +} + +// next waits for the next job, reporting false once the session is closed +// or ctx is done. +func (sess *session) next(ctx context.Context) (*job, bool) { + for { + sess.mu.Lock() + if sess.closed { + sess.mu.Unlock() + return nil, false + } + if len(sess.jobs) > 0 { + j := sess.jobs[0] + sess.jobs = sess.jobs[1:] + sess.running = j + sess.mu.Unlock() + return j, true + } + sess.mu.Unlock() + select { + case <-sess.wake: + case <-ctx.Done(): + return nil, false + } + } +} + +// finish marks j as no longer running, and reports whether it still +// needs its done, which an interrupt sends itself. +func (sess *session) finish(j *job) bool { + sess.mu.Lock() + defer sess.mu.Unlock() + if sess.running == j { + sess.running = nil + } + return !j.interrupted +} + +// interrupt marks the running job as interrupted and returns it, unless +// there's none or it isn't the one with the given id ("" for any). +func (sess *session) interrupt(id string) (j *job, mismatch bool) { + sess.mu.Lock() + defer sess.mu.Unlock() + j = sess.running + switch { + case j == nil: + return nil, false + case id != "" && id != fmt.Sprint(j.req["id"]): + return nil, true + } + j.interrupted = true + return j, false +} + +// close ends the session. Like nREPL, it interrupts the eval that's +// running, which goes on to finish the rest of its code. +func (sess *session) close() { + sess.mu.Lock() + sess.closed = true + j := sess.running + sess.mu.Unlock() + if j != nil { + signal(j.interrupts) + } + signal(sess.wake) +} + +func (sess *session) give(input string) { + sess.mu.Lock() + if input == "" { + sess.eof = true + } else { + sess.input += input + } + sess.mu.Unlock() + signal(sess.inputs) +} + +// readLine reads a line of input, asking for it with ask when there's +// none. It returns false at the end of the input, which an empty stdin +// marks. +func (sess *session) readLine(ctx context.Context, interrupts <-chan struct{}, ask func()) (string, bool, error) { + for { + // Input from before this point is in sess.input already. + select { + case <-sess.inputs: + default: + } + sess.mu.Lock() + if i := strings.IndexByte(sess.input, '\n'); i >= 0 { + line := sess.input[:i] + sess.input = sess.input[i+1:] + sess.mu.Unlock() + return line, true, nil + } + if sess.eof { + line := sess.input + sess.input, sess.eof = "", false + sess.mu.Unlock() + return line, line != "", nil + } + sess.mu.Unlock() + ask() + select { + case <-sess.inputs: + case <-interrupts: + return "", false, throwf("java.lang.InterruptedException", "") + case <-ctx.Done(): + return "", false, errStopping + } + } +}