move implementation to internal/ directory (#828)
This commit is contained in:
committed by
GitHub
parent
b9b9ad10cf
commit
908807bd59
@@ -0,0 +1,12 @@
|
||||
# setup-specific
|
||||
inventory
|
||||
*.retry
|
||||
|
||||
# generated by gen_files.sh
|
||||
files/*ssh_client_identity
|
||||
files/*ssh_client_identity.pub
|
||||
files/*.tls.*.key
|
||||
files/*.tls.*.csr
|
||||
files/*.tls.*.crt
|
||||
files/wireevaluator
|
||||
|
||||
@@ -0,0 +1,15 @@
|
||||
This directory contains very hacky test automation for wireevaluator based on nested Ansible playbooks.
|
||||
|
||||
* Copy `inventory.example` to `inventory`
|
||||
* Adjust `inventory` IP addresses as needed
|
||||
* Make sure there's an OpenSSH server running on the serve host
|
||||
* Make sure there's no firewalling whatsoever between the hosts
|
||||
* Run `GENKEYS=1 ./gen_files.sh` to re-generate self-signed TLS certs
|
||||
* Run the following command, adjusting the `wireevaluator_repeat` value to the number of times you want to repeat each test
|
||||
|
||||
```
|
||||
ansible-playbook -i inventory all.yml -e `wireevaluator_repeat=3`
|
||||
```
|
||||
|
||||
Generally, things are fine if the playbook doesn't show any panics from wireevaluator.
|
||||
|
||||
@@ -0,0 +1,17 @@
|
||||
- hosts: connect,serve
|
||||
tasks:
|
||||
|
||||
- name: "run test"
|
||||
include: internal_prepare_and_run_repeated.yml
|
||||
wireevaluator_transport: "{{config.0}}"
|
||||
wireevaluator_case: "{{config.1}}"
|
||||
wireevaluator_repeat: "{{wireevaluator_repeat}}"
|
||||
with_cartesian:
|
||||
- [ tls, ssh, tcp ]
|
||||
-
|
||||
- closewrite_server
|
||||
- closewrite_client
|
||||
- readdeadline_server
|
||||
- readdeadline_client
|
||||
loop_control:
|
||||
loop_var: config
|
||||
+55
@@ -0,0 +1,55 @@
|
||||
#!/usr/bin/env bash
|
||||
set -e
|
||||
|
||||
cd "$( dirname "${BASH_SOURCE[0]}")"
|
||||
|
||||
FILESDIR="$(pwd)"/files
|
||||
|
||||
echo "[INFO] compile binary"
|
||||
pushd .. >/dev/null
|
||||
go build -o $FILESDIR/wireevaluator
|
||||
popd >/dev/null
|
||||
|
||||
if [ "$GENKEYS" == "" ]; then
|
||||
echo "[INFO] GENKEYS environment variable not set, assumed to be valid"
|
||||
exit 0
|
||||
fi
|
||||
|
||||
echo "[INFO] gen ssh key"
|
||||
ssh-keygen -f "$FILESDIR/wireevaluator.ssh_client_identity" -t ed25519
|
||||
|
||||
echo "[INFO] gen tls keys"
|
||||
|
||||
cakey="$FILESDIR/wireevaluator.tls.ca.key"
|
||||
cacrt="$FILESDIR/wireevaluator.tls.ca.crt"
|
||||
hostprefix="$FILESDIR/wireevaluator.tls"
|
||||
|
||||
openssl genrsa -out "$cakey" 4096
|
||||
openssl req -x509 -new -nodes -key "$cakey" -sha256 -days 1 -out "$cacrt"
|
||||
|
||||
declare -a HOSTS
|
||||
HOSTS+=("theserver")
|
||||
HOSTS+=("theclient")
|
||||
|
||||
for host in "${HOSTS[@]}"; do
|
||||
key="${hostprefix}.${host}.key"
|
||||
csr="${hostprefix}.${host}.csr"
|
||||
crt="${hostprefix}.${host}.crt"
|
||||
openssl genrsa -out "$key" 2048
|
||||
|
||||
(
|
||||
echo "."
|
||||
echo "."
|
||||
echo "."
|
||||
echo "."
|
||||
echo "."
|
||||
echo $host
|
||||
echo "."
|
||||
echo "."
|
||||
echo "."
|
||||
echo "."
|
||||
) | openssl req -new -key "$key" -out "$csr"
|
||||
|
||||
openssl x509 -req -in "$csr" -CA "$cacrt" -CAkey "$cakey" -CAcreateserial -out "$crt" -days 1 -sha256
|
||||
|
||||
done
|
||||
+54
@@ -0,0 +1,54 @@
|
||||
---
|
||||
|
||||
- name: compile binary and any key files required
|
||||
local_action: command ./gen_files.sh
|
||||
|
||||
- name: Kill test binary
|
||||
shell: "killall -9 wireevaluator || true"
|
||||
- name: Deploy new binary
|
||||
copy:
|
||||
src: "files/wireevaluator"
|
||||
dest: "/opt/wireevaluator"
|
||||
mode: 0755
|
||||
|
||||
- set_fact:
|
||||
wireevaluator_connect_ip: "{{hostvars['connect'].ansible_host}}"
|
||||
wireevaluator_serve_ip: "{{hostvars['serve'].ansible_host}}"
|
||||
|
||||
- name: Deploy config
|
||||
template:
|
||||
src: "templates/{{wireevaluator_transport}}.yml.j2"
|
||||
dest: "/opt/wireevaluator.yml"
|
||||
|
||||
- name: Deploy client identity
|
||||
copy:
|
||||
src: "files/wireevaluator.{{item}}"
|
||||
dest: "/opt/wireevaluator.{{item}}"
|
||||
mode: 0400
|
||||
with_items:
|
||||
- ssh_client_identity
|
||||
- ssh_client_identity.pub
|
||||
- tls.ca.key
|
||||
- tls.ca.crt
|
||||
- tls.theserver.key
|
||||
- tls.theserver.crt
|
||||
- tls.theclient.key
|
||||
- tls.theclient.crt
|
||||
|
||||
- name: Setup server ssh client identity access
|
||||
when: inventory_hostname == "serve"
|
||||
block:
|
||||
- authorized_key:
|
||||
user: root
|
||||
state: present
|
||||
key: "{{ lookup('file', 'files/wireevaluator.ssh_client_identity.pub') }}"
|
||||
key_options: 'command="/opt/wireevaluator -mode stdinserver -config /opt/wireevaluator.yml client1"'
|
||||
- file:
|
||||
state: directory
|
||||
mode: 0700
|
||||
path: /tmp/wireevaluator_stdinserver
|
||||
|
||||
- name: repeated test
|
||||
include: internal_run_test_prepared_single.yml
|
||||
with_sequence: start=1 end={{wireevaluator_repeat}}
|
||||
|
||||
+38
@@ -0,0 +1,38 @@
|
||||
---
|
||||
|
||||
- debug:
|
||||
msg: "run test transport={{wireevaluator_transport}} case={{wireevaluator_case}} repeatedly"
|
||||
|
||||
- name: Run Server
|
||||
when: inventory_hostname == "serve"
|
||||
command: /opt/wireevaluator -config /opt/wireevaluator.yml -mode serve -testcase {{wireevaluator_case}}
|
||||
register: spawn_servers
|
||||
async: 60
|
||||
poll: 0
|
||||
|
||||
- name: Run Client
|
||||
when: inventory_hostname == "connect"
|
||||
command: /opt/wireevaluator -config /opt/wireevaluator.yml -mode connect -testcase {{wireevaluator_case}}
|
||||
register: spawn_clients
|
||||
async: 60
|
||||
poll: 0
|
||||
|
||||
- name: Wait for server shutdown
|
||||
when: inventory_hostname == "serve"
|
||||
async_status:
|
||||
jid: "{{ spawn_servers.ansible_job_id}}"
|
||||
delay: 0.5
|
||||
retries: 10
|
||||
|
||||
- name: Wait for client shutdown
|
||||
when: inventory_hostname == "connect"
|
||||
async_status:
|
||||
jid: "{{ spawn_clients.ansible_job_id}}"
|
||||
delay: 0.5
|
||||
retries: 10
|
||||
|
||||
- name: Wait for connections to die (TIME_WAIT conns)
|
||||
command: sleep 4
|
||||
changed_when: false
|
||||
|
||||
|
||||
@@ -0,0 +1,2 @@
|
||||
connect ansible_user=root ansible_host=192.168.122.128 wireevaluator_mode="connect"
|
||||
serve ansible_user=root ansible_host=192.168.122.129 wireevaluator_mode="serve"
|
||||
@@ -0,0 +1,13 @@
|
||||
connect:
|
||||
type: ssh+stdinserver
|
||||
host: {{wireevaluator_serve_ip}}
|
||||
user: root
|
||||
port: 22
|
||||
identity_file: /opt/wireevaluator.ssh_client_identity
|
||||
options: # optional, default [], `-o` arguments passed to ssh
|
||||
- "Compression=yes"
|
||||
serve:
|
||||
type: stdinserver
|
||||
client_identities:
|
||||
- "client1"
|
||||
|
||||
@@ -0,0 +1,10 @@
|
||||
connect:
|
||||
type: tcp
|
||||
address: "{{wireevaluator_serve_ip}}:8888"
|
||||
serve:
|
||||
type: tcp
|
||||
listen: ":8888"
|
||||
clients: {
|
||||
"{{wireevaluator_connect_ip}}" : "client1"
|
||||
}
|
||||
|
||||
@@ -0,0 +1,16 @@
|
||||
connect:
|
||||
type: tls
|
||||
address: "{{wireevaluator_serve_ip}}:8888"
|
||||
ca: "/opt/wireevaluator.tls.ca.crt"
|
||||
cert: "/opt/wireevaluator.tls.theclient.crt"
|
||||
key: "/opt/wireevaluator.tls.theclient.key"
|
||||
server_cn: "theserver"
|
||||
|
||||
serve:
|
||||
type: tls
|
||||
listen: ":8888"
|
||||
ca: "/opt/wireevaluator.tls.ca.crt"
|
||||
cert: "/opt/wireevaluator.tls.theserver.crt"
|
||||
key: "/opt/wireevaluator.tls.theserver.key"
|
||||
client_cns:
|
||||
- "theclient"
|
||||
@@ -0,0 +1,111 @@
|
||||
// a tool to test whether a given transport implements the timeoutconn.Wire interface
|
||||
package main
|
||||
|
||||
import (
|
||||
"context"
|
||||
"flag"
|
||||
"fmt"
|
||||
"os"
|
||||
"path"
|
||||
|
||||
netssh "github.com/problame/go-netssh"
|
||||
"github.com/zrepl/yaml-config"
|
||||
|
||||
"github.com/zrepl/zrepl/internal/config"
|
||||
"github.com/zrepl/zrepl/internal/transport"
|
||||
transportconfig "github.com/zrepl/zrepl/internal/transport/fromconfig"
|
||||
)
|
||||
|
||||
func noerror(err error) {
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
}
|
||||
|
||||
type Config struct {
|
||||
Connect config.ConnectEnum
|
||||
Serve config.ServeEnum
|
||||
}
|
||||
|
||||
var args struct {
|
||||
mode string
|
||||
configPath string
|
||||
testCase string
|
||||
}
|
||||
|
||||
var conf Config
|
||||
|
||||
type TestCase interface {
|
||||
Client(wire transport.Wire)
|
||||
Server(wire transport.Wire)
|
||||
}
|
||||
|
||||
func main() {
|
||||
flag.StringVar(&args.mode, "mode", "", "connect|serve")
|
||||
flag.StringVar(&args.configPath, "config", "", "config file path")
|
||||
flag.StringVar(&args.testCase, "testcase", "", "")
|
||||
flag.Parse()
|
||||
|
||||
bytes, err := os.ReadFile(args.configPath)
|
||||
noerror(err)
|
||||
err = yaml.UnmarshalStrict(bytes, &conf)
|
||||
noerror(err)
|
||||
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
defer cancel()
|
||||
|
||||
global := &config.Global{
|
||||
Serve: &config.GlobalServe{
|
||||
StdinServer: &config.GlobalStdinServer{
|
||||
SockDir: "/tmp/wireevaluator_stdinserver",
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
switch args.mode {
|
||||
case "connect":
|
||||
tc, err := getTestCase(args.testCase)
|
||||
noerror(err)
|
||||
connecter, err := transportconfig.ConnecterFromConfig(global, conf.Connect, config.ParseFlagsNone)
|
||||
noerror(err)
|
||||
wire, err := connecter.Connect(ctx)
|
||||
noerror(err)
|
||||
tc.Client(wire)
|
||||
case "serve":
|
||||
tc, err := getTestCase(args.testCase)
|
||||
noerror(err)
|
||||
lf, err := transportconfig.ListenerFactoryFromConfig(global, conf.Serve, config.ParseFlagsNone)
|
||||
noerror(err)
|
||||
l, err := lf()
|
||||
noerror(err)
|
||||
conn, err := l.Accept(ctx)
|
||||
noerror(err)
|
||||
tc.Server(conn)
|
||||
case "stdinserver":
|
||||
identity := flag.Arg(0)
|
||||
unixaddr := path.Join(global.Serve.StdinServer.SockDir, identity)
|
||||
err := netssh.Proxy(ctx, unixaddr)
|
||||
if err == nil {
|
||||
os.Exit(0)
|
||||
}
|
||||
panic(err)
|
||||
default:
|
||||
panic(fmt.Sprintf("unknown mode %q", args.mode))
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
func getTestCase(tcName string) (TestCase, error) {
|
||||
switch tcName {
|
||||
case "closewrite_server":
|
||||
return &CloseWrite{mode: CloseWriteServerSide}, nil
|
||||
case "closewrite_client":
|
||||
return &CloseWrite{mode: CloseWriteClientSide}, nil
|
||||
case "readdeadline_client":
|
||||
return &Deadlines{mode: DeadlineModeClientTimeout}, nil
|
||||
case "readdeadline_server":
|
||||
return &Deadlines{mode: DeadlineModeServerTimeout}, nil
|
||||
default:
|
||||
return nil, fmt.Errorf("unknown test case %q", tcName)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,106 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"io"
|
||||
"log"
|
||||
|
||||
"github.com/zrepl/zrepl/internal/transport"
|
||||
)
|
||||
|
||||
type CloseWriteMode uint
|
||||
|
||||
const (
|
||||
CloseWriteClientSide CloseWriteMode = 1 + iota
|
||||
CloseWriteServerSide
|
||||
)
|
||||
|
||||
type CloseWrite struct {
|
||||
mode CloseWriteMode
|
||||
}
|
||||
|
||||
// sent repeatedly
|
||||
var closeWriteTestSendData = bytes.Repeat([]byte{0x23, 0x42}, 1<<24)
|
||||
var closeWriteErrorMsg = []byte{0xb, 0xa, 0xd, 0xf, 0x0, 0x0, 0xd}
|
||||
|
||||
func (m CloseWrite) Client(wire transport.Wire) {
|
||||
switch m.mode {
|
||||
case CloseWriteClientSide:
|
||||
m.receiver(wire)
|
||||
case CloseWriteServerSide:
|
||||
m.sender(wire)
|
||||
default:
|
||||
panic(m.mode)
|
||||
}
|
||||
}
|
||||
|
||||
func (m CloseWrite) Server(wire transport.Wire) {
|
||||
switch m.mode {
|
||||
case CloseWriteClientSide:
|
||||
m.sender(wire)
|
||||
case CloseWriteServerSide:
|
||||
m.receiver(wire)
|
||||
default:
|
||||
panic(m.mode)
|
||||
}
|
||||
}
|
||||
|
||||
func (CloseWrite) sender(wire transport.Wire) {
|
||||
defer func() {
|
||||
closeErr := wire.Close()
|
||||
log.Printf("closeErr=%T %s", closeErr, closeErr)
|
||||
}()
|
||||
|
||||
writeDone := make(chan struct{}, 1)
|
||||
go func() {
|
||||
close(writeDone)
|
||||
for {
|
||||
_, err := wire.Write(closeWriteTestSendData)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
}
|
||||
}()
|
||||
|
||||
defer func() {
|
||||
<-writeDone
|
||||
}()
|
||||
|
||||
var respBuf bytes.Buffer
|
||||
_, err := io.Copy(&respBuf, wire)
|
||||
if err != nil {
|
||||
log.Fatalf("should have received io.EOF, which is masked by io.Copy, got: %s", err)
|
||||
}
|
||||
if !bytes.Equal(respBuf.Bytes(), closeWriteErrorMsg) {
|
||||
log.Fatalf("did not receive error message, got response with len %v:\n%v", respBuf.Len(), respBuf.Bytes())
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
func (CloseWrite) receiver(wire transport.Wire) {
|
||||
|
||||
// consume half the test data, then detect an error, send it and CloseWrite
|
||||
|
||||
r := io.LimitReader(wire, int64(5*len(closeWriteTestSendData)/3))
|
||||
_, err := io.Copy(io.Discard, r)
|
||||
noerror(err)
|
||||
|
||||
var errBuf bytes.Buffer
|
||||
errBuf.Write(closeWriteErrorMsg)
|
||||
_, err = io.Copy(wire, &errBuf)
|
||||
noerror(err)
|
||||
|
||||
err = wire.CloseWrite()
|
||||
noerror(err)
|
||||
|
||||
// drain wire, as documented in transport.Wire, this is the only way we know the client closed the conn
|
||||
_, err = io.Copy(io.Discard, wire)
|
||||
if err != nil {
|
||||
// io.Copy masks io.EOF to nil, and we expect io.EOF from the client's Close() call
|
||||
log.Panicf("unexpected error returned from reading conn: %s", err)
|
||||
}
|
||||
|
||||
closeErr := wire.Close()
|
||||
log.Printf("closeErr=%T %s", closeErr, closeErr)
|
||||
|
||||
}
|
||||
@@ -0,0 +1,138 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"fmt"
|
||||
"io"
|
||||
"log"
|
||||
"net"
|
||||
"time"
|
||||
|
||||
"github.com/zrepl/zrepl/internal/transport"
|
||||
)
|
||||
|
||||
type DeadlineMode uint
|
||||
|
||||
const (
|
||||
DeadlineModeClientTimeout DeadlineMode = 1 + iota
|
||||
DeadlineModeServerTimeout
|
||||
)
|
||||
|
||||
type Deadlines struct {
|
||||
mode DeadlineMode
|
||||
}
|
||||
|
||||
func (d Deadlines) Client(wire transport.Wire) {
|
||||
switch d.mode {
|
||||
case DeadlineModeClientTimeout:
|
||||
d.sleepThenSend(wire)
|
||||
case DeadlineModeServerTimeout:
|
||||
d.sendThenRead(wire)
|
||||
default:
|
||||
panic(d.mode)
|
||||
}
|
||||
}
|
||||
|
||||
func (d Deadlines) Server(wire transport.Wire) {
|
||||
switch d.mode {
|
||||
case DeadlineModeClientTimeout:
|
||||
d.sendThenRead(wire)
|
||||
case DeadlineModeServerTimeout:
|
||||
d.sleepThenSend(wire)
|
||||
default:
|
||||
panic(d.mode)
|
||||
}
|
||||
}
|
||||
|
||||
var deadlinesTimeout = 1 * time.Second
|
||||
|
||||
func (d Deadlines) sleepThenSend(wire transport.Wire) {
|
||||
defer wire.Close()
|
||||
|
||||
log.Print("sleepThenSend")
|
||||
|
||||
// exceed timeout of peer (do not respond to their hi msg)
|
||||
time.Sleep(3 * deadlinesTimeout)
|
||||
// expect that the client has hung up on us by now
|
||||
err := d.sendMsg(wire, "hi")
|
||||
log.Printf("err=%s", err)
|
||||
log.Printf("err=%#v", err)
|
||||
if err == nil {
|
||||
log.Panic("no error")
|
||||
}
|
||||
if _, ok := err.(net.Error); !ok {
|
||||
log.Panic("not a net error")
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
func (d Deadlines) sendThenRead(wire transport.Wire) {
|
||||
|
||||
log.Print("sendThenRead")
|
||||
|
||||
err := d.sendMsg(wire, "hi")
|
||||
noerror(err)
|
||||
|
||||
err = wire.SetReadDeadline(time.Now().Add(deadlinesTimeout))
|
||||
noerror(err)
|
||||
|
||||
m, err := d.recvMsg(wire)
|
||||
log.Printf("m=%q", m)
|
||||
log.Printf("err=%s", err)
|
||||
log.Printf("err=%#v", err)
|
||||
|
||||
// close asap so that the peer get's a 'connection reset by peer' error or similar
|
||||
closeErr := wire.Close()
|
||||
if closeErr != nil {
|
||||
panic(closeErr)
|
||||
}
|
||||
|
||||
var neterr net.Error
|
||||
var ok bool
|
||||
if err == nil {
|
||||
goto unexpErr // works for nil, too
|
||||
}
|
||||
neterr, ok = err.(net.Error)
|
||||
if !ok {
|
||||
log.Println("not a net error")
|
||||
goto unexpErr
|
||||
}
|
||||
if !neterr.Timeout() {
|
||||
log.Println("not a timeout")
|
||||
}
|
||||
|
||||
return
|
||||
|
||||
unexpErr:
|
||||
panic(fmt.Sprintf("sendThenRead: client should have hung up but got error %T %s", err, err))
|
||||
}
|
||||
|
||||
const deadlinesMsgLen = 40
|
||||
|
||||
func (d Deadlines) sendMsg(wire transport.Wire, msg string) error {
|
||||
if len(msg) > deadlinesMsgLen {
|
||||
panic(len(msg))
|
||||
}
|
||||
var buf [deadlinesMsgLen]byte
|
||||
copy(buf[:], []byte(msg))
|
||||
n, err := wire.Write(buf[:])
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if n != len(buf) {
|
||||
panic("short write not allowed")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (d Deadlines) recvMsg(wire transport.Wire) (string, error) {
|
||||
|
||||
var buf bytes.Buffer
|
||||
r := io.LimitReader(wire, deadlinesMsgLen)
|
||||
_, err := io.Copy(&buf, r)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return buf.String(), nil
|
||||
|
||||
}
|
||||
@@ -0,0 +1,210 @@
|
||||
// package timeoutconn wraps a Wire to provide idle timeouts
|
||||
// based on Set{Read,Write}Deadline.
|
||||
// Additionally, it exports abstractions for vectored I/O.
|
||||
package timeoutconn
|
||||
|
||||
// NOTE
|
||||
// Readv and Writev are not split-off into a separate package
|
||||
// because we use raw syscalls, bypassing Conn's Read / Write methods.
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"io"
|
||||
"net"
|
||||
"sync"
|
||||
"syscall"
|
||||
"time"
|
||||
)
|
||||
|
||||
type Wire interface {
|
||||
net.Conn
|
||||
// A call to CloseWrite indicates that no further Write calls will be made to Wire.
|
||||
// The implementation must return an error in case of Write calls after CloseWrite.
|
||||
// On the peer's side, after it read all data written to Wire prior to the call to
|
||||
// CloseWrite on our side, the peer's Read calls must return io.EOF.
|
||||
// CloseWrite must not affect the read-direction of Wire: specifically, the
|
||||
// peer must continue to be able to send, and our side must continue be
|
||||
// able to receive data over Wire.
|
||||
//
|
||||
// Note that CloseWrite may (and most likely will) return sooner than the
|
||||
// peer having received all data written to Wire prior to CloseWrite.
|
||||
// Note further that buffering happening in the network stacks on either side
|
||||
// mandates an explicit acknowledgement from the peer that the connection may
|
||||
// be fully shut down: If we call Close without such acknowledgement, any data
|
||||
// from peer to us that was already in flight may cause connection resets to
|
||||
// be sent from us to the peer via the specific transport protocol. Those
|
||||
// resets (e.g. RST frames) may erase all connection context on the peer,
|
||||
// including data in its receive buffers. Thus, those resets are in race with
|
||||
// a) transmission of data written prior to CloseWrite and
|
||||
// b) the peer application reading from those buffers.
|
||||
//
|
||||
// The WaitForPeerClose method can be used to wait for connection termination,
|
||||
// iff the implementation supports it. If it does not, the only reliable way
|
||||
// to wait for a peer to have read all data from Wire (until io.EOF), is to
|
||||
// expect it to close the wire at that point as well, and to drain Wire until
|
||||
// we also read io.EOF.
|
||||
CloseWrite() error
|
||||
|
||||
// Wait for the peer to close the connection.
|
||||
// No data that could otherwise be Read is lost as a consequence of this call.
|
||||
// The use case for this API is abortive connection shutdown.
|
||||
// To provide any value over draining Wire using io.Read, an implementation
|
||||
// will likely use out-of-band messaging mechanisms.
|
||||
// TODO WaitForPeerClose() (supported bool, err error)
|
||||
}
|
||||
|
||||
type Conn struct {
|
||||
// immutable state
|
||||
|
||||
Wire
|
||||
idleTimeout time.Duration
|
||||
|
||||
// mutable state (protected by mtx)
|
||||
|
||||
mtx sync.RWMutex
|
||||
renewDeadlinesDisabled bool
|
||||
}
|
||||
|
||||
func Wrap(conn Wire, idleTimeout time.Duration) *Conn {
|
||||
return &Conn{
|
||||
Wire: conn,
|
||||
idleTimeout: idleTimeout,
|
||||
}
|
||||
}
|
||||
|
||||
// DisableTimeouts disables the idle timeout behavior provided by this package.
|
||||
// Existing deadlines are cleared iff the call is the first call to this method
|
||||
// or if the previous call produced an error.
|
||||
func (c *Conn) DisableTimeouts() error {
|
||||
c.mtx.Lock()
|
||||
defer c.mtx.Unlock()
|
||||
if c.renewDeadlinesDisabled {
|
||||
return nil
|
||||
}
|
||||
err := c.SetDeadline(time.Time{})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
c.renewDeadlinesDisabled = true
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *Conn) renewReadDeadline() error {
|
||||
c.mtx.RLock()
|
||||
defer c.mtx.RUnlock()
|
||||
if c.renewDeadlinesDisabled {
|
||||
return nil
|
||||
}
|
||||
return c.SetReadDeadline(time.Now().Add(c.idleTimeout))
|
||||
}
|
||||
|
||||
func (c *Conn) RenewWriteDeadline() error {
|
||||
c.mtx.RLock()
|
||||
defer c.mtx.RUnlock()
|
||||
if c.renewDeadlinesDisabled {
|
||||
return nil
|
||||
}
|
||||
return c.SetWriteDeadline(time.Now().Add(c.idleTimeout))
|
||||
}
|
||||
|
||||
func (c *Conn) Read(p []byte) (n int, _ error) {
|
||||
n = 0
|
||||
restart:
|
||||
if err := c.renewReadDeadline(); err != nil {
|
||||
return n, err
|
||||
}
|
||||
var nCurRead int
|
||||
nCurRead, err := c.Wire.Read(p[n:])
|
||||
n += nCurRead
|
||||
if netErr, ok := err.(net.Error); ok && netErr.Timeout() && nCurRead > 0 {
|
||||
goto restart
|
||||
}
|
||||
return n, err
|
||||
}
|
||||
|
||||
func (c *Conn) Write(p []byte) (n int, _ error) {
|
||||
n = 0
|
||||
restart:
|
||||
if err := c.RenewWriteDeadline(); err != nil {
|
||||
return n, err
|
||||
}
|
||||
var nCurWrite int
|
||||
nCurWrite, err := c.Wire.Write(p[n:])
|
||||
n += nCurWrite
|
||||
if netErr, ok := err.(net.Error); ok && netErr.Timeout() && nCurWrite > 0 {
|
||||
goto restart
|
||||
}
|
||||
return n, err
|
||||
}
|
||||
|
||||
// Writes the given buffers to Conn, following the semantics of io.Copy,
|
||||
// but is guaranteed to use the writev system call if the wrapped Wire
|
||||
// support it.
|
||||
// Note the Conn does not support writev through io.Copy(aConn, aNetBuffers).
|
||||
func (c *Conn) WritevFull(bufs net.Buffers) (n int64, _ error) {
|
||||
n = 0
|
||||
restart:
|
||||
if err := c.RenewWriteDeadline(); err != nil {
|
||||
return n, err
|
||||
}
|
||||
var nCurWrite int64
|
||||
nCurWrite, err := io.Copy(c.Wire, &bufs)
|
||||
n += nCurWrite
|
||||
if netErr, ok := err.(net.Error); ok && netErr.Timeout() && nCurWrite > 0 {
|
||||
goto restart
|
||||
}
|
||||
return n, err
|
||||
}
|
||||
|
||||
var SyscallConnNotSupported = errors.New("SyscallConn not supported")
|
||||
|
||||
// The interface that must be implemented for vectored I/O support.
|
||||
// If the wrapped Wire does not implement it, a less efficient
|
||||
// fallback implementation is used.
|
||||
// Rest assured that Go's *net.TCPConn implements this interface.
|
||||
type SyscallConner interface {
|
||||
// The sentinel error value SyscallConnNotSupported can be returned
|
||||
// if the support for SyscallConn depends on runtime conditions and
|
||||
// that runtime condition is not met.
|
||||
SyscallConn() (syscall.RawConn, error)
|
||||
}
|
||||
|
||||
var _ SyscallConner = (*net.TCPConn)(nil)
|
||||
|
||||
// Reads the given buffers full:
|
||||
// Think of io.ReadvFull, but for net.Buffers + using the readv syscall.
|
||||
//
|
||||
// If the underlying Wire is not a SyscallConner, a fallback
|
||||
// implementation based on repeated Conn.Read invocations is used.
|
||||
//
|
||||
// If the connection returned io.EOF, the number of bytes written until
|
||||
// then + io.EOF is returned. This behavior is different to io.ReadFull
|
||||
// which returns io.ErrUnexpectedEOF.
|
||||
func (c *Conn) ReadvFull(buffers net.Buffers) (n int64, err error) {
|
||||
return c.readv(buffers)
|
||||
}
|
||||
|
||||
// invoked by c.readv if readv system call cannot be used
|
||||
func (c *Conn) readvFallback(nbuffers net.Buffers) (n int64, err error) {
|
||||
buffers := [][]byte(nbuffers)
|
||||
for i := range buffers {
|
||||
curBuf := buffers[i]
|
||||
inner:
|
||||
for len(curBuf) > 0 {
|
||||
if err := c.renewReadDeadline(); err != nil {
|
||||
return n, err
|
||||
}
|
||||
var oneN int
|
||||
oneN, err = c.Read(curBuf[:]) // WE WANT NO SHADOWING
|
||||
curBuf = curBuf[oneN:]
|
||||
n += int64(oneN)
|
||||
if err != nil {
|
||||
if netErr, ok := err.(net.Error); ok && netErr.Timeout() && oneN > 0 {
|
||||
continue inner
|
||||
}
|
||||
return n, err
|
||||
}
|
||||
}
|
||||
}
|
||||
return n, nil
|
||||
}
|
||||
@@ -0,0 +1,12 @@
|
||||
//go:build illumos || solaris
|
||||
// +build illumos solaris
|
||||
|
||||
package timeoutconn
|
||||
|
||||
import "net"
|
||||
|
||||
func (c *Conn) readv(buffers net.Buffers) (n int64, err error) {
|
||||
// Go does not expose the SYS_READV symbol for Solaris / Illumos - do they have it?
|
||||
// Anyhow, use the fallback
|
||||
return c.readvFallback(buffers)
|
||||
}
|
||||
@@ -0,0 +1,195 @@
|
||||
package timeoutconn
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"io"
|
||||
"net"
|
||||
"sync"
|
||||
"syscall"
|
||||
"testing"
|
||||
"time"
|
||||
"unsafe"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/zrepl/zrepl/internal/util/socketpair"
|
||||
"github.com/zrepl/zrepl/internal/util/zreplcircleci"
|
||||
)
|
||||
|
||||
func TestReadTimeout(t *testing.T) {
|
||||
|
||||
a, b, err := socketpair.SocketPair()
|
||||
require.NoError(t, err)
|
||||
defer a.Close()
|
||||
defer b.Close()
|
||||
|
||||
var wg sync.WaitGroup
|
||||
wg.Add(2)
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
var buf bytes.Buffer
|
||||
buf.WriteString("tooktoolong")
|
||||
time.Sleep(500 * time.Millisecond)
|
||||
_, err := io.Copy(a, &buf)
|
||||
require.NoError(t, err)
|
||||
}()
|
||||
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
conn := Wrap(b, 100*time.Millisecond)
|
||||
buf := [4]byte{} // shorter than message put on wire
|
||||
n, err := conn.Read(buf[:])
|
||||
assert.Equal(t, 0, n)
|
||||
assert.Error(t, err)
|
||||
netErr, ok := err.(net.Error)
|
||||
require.True(t, ok)
|
||||
assert.True(t, netErr.Timeout())
|
||||
}()
|
||||
|
||||
wg.Wait()
|
||||
}
|
||||
|
||||
type writeBlockConn struct {
|
||||
net.Conn
|
||||
blockTime time.Duration
|
||||
}
|
||||
|
||||
func (c writeBlockConn) Write(p []byte) (int, error) {
|
||||
time.Sleep(c.blockTime)
|
||||
return c.Conn.Write(p)
|
||||
}
|
||||
|
||||
func (c writeBlockConn) CloseWrite() error {
|
||||
return c.Conn.Close()
|
||||
}
|
||||
|
||||
func TestWriteTimeout(t *testing.T) {
|
||||
a, b, err := socketpair.SocketPair()
|
||||
require.NoError(t, err)
|
||||
defer a.Close()
|
||||
defer b.Close()
|
||||
var buf bytes.Buffer
|
||||
buf.WriteString("message")
|
||||
blockConn := writeBlockConn{a, 500 * time.Millisecond}
|
||||
conn := Wrap(blockConn, 100*time.Millisecond)
|
||||
n, err := conn.Write(buf.Bytes())
|
||||
assert.Equal(t, 0, n)
|
||||
assert.Error(t, err)
|
||||
netErr, ok := err.(net.Error)
|
||||
require.True(t, ok)
|
||||
assert.True(t, netErr.Timeout())
|
||||
}
|
||||
|
||||
func TestNoPartialReadsDueToDeadline(t *testing.T) {
|
||||
zreplcircleci.SkipOnCircleCI(t, "needs predictable low scheduling latency")
|
||||
|
||||
a, b, err := socketpair.SocketPair()
|
||||
require.NoError(t, err)
|
||||
defer a.Close()
|
||||
defer b.Close()
|
||||
|
||||
var wg sync.WaitGroup
|
||||
wg.Add(2)
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
a.Write([]byte{1, 2, 3, 4, 5})
|
||||
// sleep to provoke a partial read in the consumer goroutine
|
||||
time.Sleep(50 * time.Millisecond)
|
||||
a.Write([]byte{6, 7, 8, 9, 10})
|
||||
}()
|
||||
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
bc := Wrap(b, 100*time.Millisecond)
|
||||
var buf bytes.Buffer
|
||||
beginRead := time.Now()
|
||||
// io.Copy will encounter a partial read, then wait ~50ms until the other 5 bytes are written
|
||||
// It is still going to fail with deadline err because it expects EOF
|
||||
n, err := io.Copy(&buf, bc)
|
||||
readDuration := time.Since(beginRead)
|
||||
t.Logf("read duration=%s", readDuration)
|
||||
t.Logf("recv done n=%v err=%v", n, err)
|
||||
t.Logf("buf=%v", buf.Bytes())
|
||||
neterr, ok := err.(net.Error)
|
||||
require.True(t, ok)
|
||||
assert.True(t, neterr.Timeout())
|
||||
|
||||
assert.Equal(t, int64(10), n)
|
||||
assert.Equal(t, []byte{1, 2, 3, 4, 5, 6, 7, 8, 9, 10}, buf.Bytes())
|
||||
// 50ms for the second read, 100ms after that one for the deadline
|
||||
// allow for some jitter
|
||||
assert.True(t, readDuration > 140*time.Millisecond)
|
||||
assert.True(t, readDuration < 200*time.Millisecond)
|
||||
}()
|
||||
|
||||
wg.Wait()
|
||||
}
|
||||
|
||||
type partialWriteMockConn struct {
|
||||
net.Conn // to satisfy interface
|
||||
buf bytes.Buffer
|
||||
writeDuration time.Duration
|
||||
returnAfterBytesWritten int
|
||||
}
|
||||
|
||||
func newPartialWriteMockConn(writeDuration time.Duration, returnAfterBytesWritten int) *partialWriteMockConn {
|
||||
return &partialWriteMockConn{
|
||||
writeDuration: writeDuration,
|
||||
returnAfterBytesWritten: returnAfterBytesWritten,
|
||||
}
|
||||
}
|
||||
|
||||
func (c *partialWriteMockConn) Write(p []byte) (int, error) {
|
||||
time.Sleep(c.writeDuration)
|
||||
consumeBytes := len(p)
|
||||
if consumeBytes > c.returnAfterBytesWritten {
|
||||
consumeBytes = c.returnAfterBytesWritten
|
||||
}
|
||||
n, err := c.buf.Write(p[0:consumeBytes])
|
||||
if err != nil || n != consumeBytes {
|
||||
panic("bytes.Buffer behaves unexpectedly")
|
||||
}
|
||||
return n, nil
|
||||
}
|
||||
|
||||
func TestPartialWriteMockConn(t *testing.T) {
|
||||
zreplcircleci.SkipOnCircleCI(t, "because it relies on scheduler responsiveness < 50ms")
|
||||
mc := newPartialWriteMockConn(100*time.Millisecond, 5)
|
||||
buf := []byte{1, 2, 3, 4, 5, 6, 7, 8, 9, 10}
|
||||
begin := time.Now()
|
||||
n, err := mc.Write(buf[:])
|
||||
duration := time.Since(begin)
|
||||
assert.NoError(t, err)
|
||||
assert.Equal(t, 5, n)
|
||||
assert.True(t, duration > 100*time.Millisecond)
|
||||
assert.True(t, duration < 150*time.Millisecond)
|
||||
}
|
||||
|
||||
func TestNoPartialWritesDueToDeadline(t *testing.T) {
|
||||
a, b, err := socketpair.SocketPair()
|
||||
require.NoError(t, err)
|
||||
defer a.Close()
|
||||
defer b.Close()
|
||||
var buf bytes.Buffer
|
||||
buf.WriteString("message")
|
||||
blockConn := writeBlockConn{a, 150 * time.Millisecond}
|
||||
conn := Wrap(blockConn, 100*time.Millisecond)
|
||||
n, err := conn.Write(buf.Bytes())
|
||||
assert.Equal(t, 0, n)
|
||||
assert.Error(t, err)
|
||||
netErr, ok := err.(net.Error)
|
||||
require.True(t, ok)
|
||||
assert.True(t, netErr.Timeout())
|
||||
}
|
||||
|
||||
func TestIovecLenFieldIsMachineUint(t *testing.T) {
|
||||
iov := syscall.Iovec{}
|
||||
_ = iov // make linter happy (unsafe.Sizeof not recognized as usage)
|
||||
size_t := unsafe.Sizeof(iov.Len)
|
||||
if size_t != unsafe.Sizeof(uint(23)) {
|
||||
t.Fatalf("expecting (struct iov)->Len to be sizeof(uint)")
|
||||
}
|
||||
// ssize_t is defined to be the signed version of size_t,
|
||||
// so we know sizeof(ssize_t) == sizeof(int)
|
||||
}
|
||||
@@ -0,0 +1,132 @@
|
||||
//go:build !illumos && !solaris
|
||||
// +build !illumos,!solaris
|
||||
|
||||
package timeoutconn
|
||||
|
||||
import (
|
||||
"io"
|
||||
"net"
|
||||
"syscall"
|
||||
"unsafe"
|
||||
)
|
||||
|
||||
func buildIovecs(buffers net.Buffers) (totalLen int64, vecs []syscall.Iovec) {
|
||||
vecs = make([]syscall.Iovec, 0, len(buffers))
|
||||
for i := range buffers {
|
||||
totalLen += int64(len(buffers[i]))
|
||||
if len(buffers[i]) == 0 {
|
||||
continue
|
||||
}
|
||||
|
||||
v := syscall.Iovec{
|
||||
Base: &buffers[i][0],
|
||||
}
|
||||
// syscall.Iovec.Len has platform-dependent size, thus use SetLen
|
||||
v.SetLen(len(buffers[i]))
|
||||
|
||||
vecs = append(vecs, v)
|
||||
}
|
||||
return totalLen, vecs
|
||||
}
|
||||
|
||||
func (c *Conn) readv(buffers net.Buffers) (n int64, err error) {
|
||||
scc, ok := c.Wire.(SyscallConner)
|
||||
if !ok {
|
||||
return c.readvFallback(buffers)
|
||||
}
|
||||
rawConn, err := scc.SyscallConn()
|
||||
if err == SyscallConnNotSupported {
|
||||
return c.readvFallback(buffers)
|
||||
}
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
|
||||
_, iovecs := buildIovecs(buffers)
|
||||
|
||||
for len(iovecs) > 0 {
|
||||
if err := c.renewReadDeadline(); err != nil {
|
||||
return n, err
|
||||
}
|
||||
oneN, oneErr := c.doOneReadv(rawConn, &iovecs)
|
||||
n += oneN
|
||||
if netErr, ok := oneErr.(net.Error); ok && netErr.Timeout() && oneN > 0 { // TODO likely not working
|
||||
continue
|
||||
} else if oneErr == nil && oneN > 0 {
|
||||
continue
|
||||
} else {
|
||||
return n, oneErr
|
||||
}
|
||||
}
|
||||
return n, nil
|
||||
}
|
||||
|
||||
func (c *Conn) doOneReadv(rawConn syscall.RawConn, iovecs *[]syscall.Iovec) (n int64, err error) {
|
||||
rawReadErr := rawConn.Read(func(fd uintptr) (done bool) {
|
||||
// iovecs, n and err must not be shadowed!
|
||||
|
||||
// NOTE: unsafe.Pointer safety rules
|
||||
// https://tip.golang.org/pkg/unsafe/#Pointer
|
||||
//
|
||||
// (4) Conversion of a Pointer to a uintptr when calling syscall.Syscall.
|
||||
// ...
|
||||
// uintptr() conversions must appear within the syscall.Syscall argument list.
|
||||
// (even though we are not the escape analysis Likely not )
|
||||
thisReadN, _, errno := syscall.Syscall(
|
||||
syscall.SYS_READV,
|
||||
fd,
|
||||
uintptr(unsafe.Pointer(&(*iovecs)[0])),
|
||||
uintptr(len(*iovecs)),
|
||||
)
|
||||
if thisReadN == ^uintptr(0) {
|
||||
if errno == syscall.EAGAIN {
|
||||
return false
|
||||
}
|
||||
err = syscall.Errno(errno)
|
||||
return true
|
||||
}
|
||||
if int(thisReadN) < 0 {
|
||||
panic("unexpected return value")
|
||||
}
|
||||
n += int64(thisReadN) // TODO check overflow
|
||||
|
||||
// shift iovecs forward
|
||||
for left := int(thisReadN); left > 0; {
|
||||
// conversion to uint does not change value, see TestIovecLenFieldIsMachineUint, and left > 0
|
||||
thisIovecConsumedCompletely := uint((*iovecs)[0].Len) <= uint(left)
|
||||
if thisIovecConsumedCompletely {
|
||||
// Update left, cannot go below 0 due to
|
||||
// a) definition of thisIovecConsumedCompletely
|
||||
// b) left > 0 due to loop invariant
|
||||
// Converting .Len to int64 is thus also safe now, because it is < left < INT_MAX
|
||||
left -= int((*iovecs)[0].Len)
|
||||
*iovecs = (*iovecs)[1:]
|
||||
} else {
|
||||
// trim this iovec to remaining length
|
||||
|
||||
// NOTE: unsafe.Pointer safety rules
|
||||
// https://tip.golang.org/pkg/unsafe/#Pointer
|
||||
// (3) Conversion of a Pointer to a uintptr and back, with arithmetic.
|
||||
// ...
|
||||
// Note that both conversions must appear in the same expression,
|
||||
// with only the intervening arithmetic between them:
|
||||
(*iovecs)[0].Base = (*byte)(unsafe.Pointer(uintptr(unsafe.Pointer((*iovecs)[0].Base)) + uintptr(left)))
|
||||
curVecNewLength := uint((*iovecs)[0].Len) - uint(left) // casts to uint do not change value
|
||||
(*iovecs)[0].SetLen(int(curVecNewLength)) // int and uint have the same size, no change of value
|
||||
|
||||
break // inner
|
||||
}
|
||||
}
|
||||
if thisReadN == 0 {
|
||||
err = io.EOF
|
||||
return true
|
||||
}
|
||||
return true
|
||||
})
|
||||
|
||||
if rawReadErr != nil {
|
||||
err = rawReadErr
|
||||
}
|
||||
|
||||
return n, err
|
||||
}
|
||||
Reference in New Issue
Block a user