mirror of
https://github.com/Harrison-Wells-Cyber/PS-Proxy
synced 2026-07-26 08:06:34 +00:00
269 lines
7.1 KiB
Go
269 lines
7.1 KiB
Go
package main
|
|
|
|
import (
|
|
"bufio"
|
|
"crypto/hmac"
|
|
"crypto/rand"
|
|
"crypto/rsa"
|
|
"crypto/sha1"
|
|
"crypto/sha256"
|
|
"crypto/x509"
|
|
"encoding/base64"
|
|
"encoding/binary"
|
|
"io"
|
|
"net"
|
|
"strings"
|
|
"testing"
|
|
"time"
|
|
|
|
"github.com/psproxy/psproxy/internal/protocol"
|
|
"github.com/psproxy/psproxy/internal/staging"
|
|
)
|
|
|
|
func TestTunnelServerMaxStreams(t *testing.T) {
|
|
server := NewTunnelServer(staging.NewStore(time.Minute), 1)
|
|
agent, peer := net.Pipe()
|
|
defer peer.Close()
|
|
sess := NewAgentSession(agent, bufio.NewReader(agent), server.maxStreams)
|
|
defer sess.Close()
|
|
server.SetSession(sess)
|
|
|
|
local1, remote1 := net.Pipe()
|
|
defer remote1.Close()
|
|
defer local1.Close()
|
|
if err := sess.AttachLocal(1, local1); err != nil {
|
|
t.Fatalf("first stream should attach: %v", err)
|
|
}
|
|
|
|
local2, remote2 := net.Pipe()
|
|
defer remote2.Close()
|
|
defer local2.Close()
|
|
if err := sess.AttachLocal(2, local2); err == nil {
|
|
t.Fatal("second stream should be rejected at max stream limit")
|
|
}
|
|
}
|
|
|
|
func TestLocalStreamBackpressureClosesConnection(t *testing.T) {
|
|
local, remote := net.Pipe()
|
|
defer remote.Close()
|
|
ls := newLocalStream(local)
|
|
defer ls.close()
|
|
|
|
failed := false
|
|
for i := 0; i < cap(ls.ch)+10; i++ {
|
|
if !ls.enqueue([]byte("x")) {
|
|
failed = true
|
|
break
|
|
}
|
|
}
|
|
if !failed {
|
|
t.Fatal("enqueue should eventually fail when the local write queue is full")
|
|
}
|
|
}
|
|
|
|
func TestTunnelServerQueryDNS(t *testing.T) {
|
|
server := NewTunnelServer(staging.NewStore(time.Minute), 2)
|
|
agent, peer := net.Pipe()
|
|
defer peer.Close()
|
|
sess := NewAgentSession(agent, bufio.NewReader(agent), server.maxStreams)
|
|
defer sess.Close()
|
|
server.SetSession(sess)
|
|
go sess.Run()
|
|
|
|
go func() {
|
|
f, err := protocol.ReadFrame(peer)
|
|
if err != nil {
|
|
return
|
|
}
|
|
_ = protocol.WriteFrame(peer, protocol.Frame{StreamID: f.StreamID, Type: protocol.FrameDNSReply, Payload: []byte("dns-response")})
|
|
}()
|
|
|
|
resp, err := server.QueryDNS([]byte("dns-query"))
|
|
if err != nil {
|
|
t.Fatalf("dns query failed: %v", err)
|
|
}
|
|
if string(resp) != "dns-response" {
|
|
t.Fatalf("unexpected dns response: %q", resp)
|
|
}
|
|
}
|
|
|
|
func TestServeDNSTCPForwardsLengthPrefixedQuery(t *testing.T) {
|
|
server := NewTunnelServer(staging.NewStore(time.Minute), 2)
|
|
agent, peer := net.Pipe()
|
|
defer peer.Close()
|
|
sess := NewAgentSession(agent, bufio.NewReader(agent), server.maxStreams)
|
|
defer sess.Close()
|
|
server.SetSession(sess)
|
|
go sess.Run()
|
|
|
|
ln, err := net.Listen("tcp", "127.0.0.1:0")
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
defer ln.Close()
|
|
go serveDNSTCP(ln, server)
|
|
|
|
go func() {
|
|
f, err := protocol.ReadFrame(peer)
|
|
if err != nil {
|
|
return
|
|
}
|
|
if string(f.Payload) != "dns-query" {
|
|
return
|
|
}
|
|
_ = protocol.WriteFrame(peer, protocol.Frame{StreamID: f.StreamID, Type: protocol.FrameDNSReply, Payload: []byte("dns-response")})
|
|
}()
|
|
|
|
conn, err := net.Dial("tcp", ln.Addr().String())
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
defer conn.Close()
|
|
var lenb [2]byte
|
|
binary.BigEndian.PutUint16(lenb[:], uint16(len("dns-query")))
|
|
if _, err := conn.Write(append(lenb[:], []byte("dns-query")...)); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if _, err := io.ReadFull(conn, lenb[:]); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
resp := make([]byte, binary.BigEndian.Uint16(lenb[:]))
|
|
if _, err := io.ReadFull(conn, resp); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if string(resp) != "dns-response" {
|
|
t.Fatalf("unexpected dns tcp response: %q", resp)
|
|
}
|
|
}
|
|
|
|
func TestValidateRoutes(t *testing.T) {
|
|
if err := validateRoutes([]string{"10.0.0.0/24", "192.168.1.10/32"}); err != nil {
|
|
t.Fatalf("valid routes rejected: %v", err)
|
|
}
|
|
if err := validateRoutes([]string{"not-a-cidr"}); err == nil {
|
|
t.Fatal("invalid CIDR should be rejected")
|
|
}
|
|
}
|
|
|
|
func TestTunnelServerQueryDNSEmptyResponseFails(t *testing.T) {
|
|
server := NewTunnelServer(staging.NewStore(time.Minute), 2)
|
|
agent, peer := net.Pipe()
|
|
defer peer.Close()
|
|
sess := NewAgentSession(agent, bufio.NewReader(agent), server.maxStreams)
|
|
defer sess.Close()
|
|
server.SetSession(sess)
|
|
go sess.Run()
|
|
|
|
go func() {
|
|
f, err := protocol.ReadFrame(peer)
|
|
if err != nil {
|
|
return
|
|
}
|
|
_ = protocol.WriteFrame(peer, protocol.Frame{StreamID: f.StreamID, Type: protocol.FrameDNSReply})
|
|
}()
|
|
|
|
if _, err := server.QueryDNS([]byte("dns-query")); err == nil {
|
|
t.Fatal("empty DNS replies should fail")
|
|
}
|
|
}
|
|
|
|
func TestSingleListenerAcceptReturnsEOFAfterFirstConn(t *testing.T) {
|
|
c1, c2 := net.Pipe()
|
|
defer c1.Close()
|
|
defer c2.Close()
|
|
ln := &singleListener{conn: c1, done: make(chan struct{})}
|
|
got, err := ln.Accept()
|
|
if err != nil {
|
|
t.Fatalf("first accept failed: %v", err)
|
|
}
|
|
if got != c1 {
|
|
t.Fatal("first accept returned unexpected connection")
|
|
}
|
|
if _, err := ln.Accept(); err == nil {
|
|
t.Fatal("second accept should return EOF")
|
|
}
|
|
}
|
|
func TestDecryptSessionSecretAcceptsAgentSHA1OAEP(t *testing.T) {
|
|
key, err := rsa.GenerateKey(rand.Reader, 3072)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
secret := []byte("0123456789abcdef0123456789abcdef")
|
|
enc, err := rsa.EncryptOAEP(sha1.New(), rand.Reader, &key.PublicKey, secret, nil)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
got, err := decryptSessionSecret(key, enc)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if string(got) != string(secret) {
|
|
t.Fatalf("got %q want %q", got, secret)
|
|
}
|
|
}
|
|
|
|
func TestServerSecureHandshakeEncryptedFrame(t *testing.T) {
|
|
key, err := rsa.GenerateKey(rand.Reader, 3072)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
store := staging.NewStore(time.Minute)
|
|
sess, err := store.Create("c2.example.com", 443, "server-key", "")
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
server := NewTunnelServer(store, 2)
|
|
server.identityKey = key
|
|
srv, cli := net.Pipe()
|
|
defer cli.Close()
|
|
go handleAgent(srv, bufio.NewReader(srv), server)
|
|
if _, err := cli.Write([]byte("HELLO ")); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
secret := []byte("0123456789abcdef0123456789abcdef")
|
|
nonce := []byte("abcdef0123456789abcdef0123456789")
|
|
enc, err := rsa.EncryptOAEP(sha256.New(), rand.Reader, &key.PublicKey, secret, []byte("PS-Proxy PSP1 session"))
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
helloTail := base64.RawURLEncoding.EncodeToString(enc) + " " + base64.RawURLEncoding.EncodeToString(nonce) + "\n"
|
|
if _, err := cli.Write([]byte(helloTail)); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
br := bufio.NewReader(cli)
|
|
line, err := br.ReadString('\n')
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
parts := strings.Fields(strings.TrimSpace(line))
|
|
if len(parts) != 3 || parts[0] != "PROOF" {
|
|
t.Fatalf("bad proof line: %q", line)
|
|
}
|
|
serverNonce, _ := base64.RawURLEncoding.DecodeString(parts[1])
|
|
proof, _ := base64.RawURLEncoding.DecodeString(parts[2])
|
|
pubDER, _ := x509.MarshalPKIXPublicKey(&key.PublicKey)
|
|
mac := hmac.New(sha256.New, secret)
|
|
mac.Write([]byte(protocol.Magic))
|
|
mac.Write([]byte("HELLO " + helloTail))
|
|
mac.Write(serverNonce)
|
|
mac.Write(nonce)
|
|
mac.Write(pubDER)
|
|
if !hmac.Equal(proof, mac.Sum(nil)) {
|
|
t.Fatal("proof mismatch")
|
|
}
|
|
codec, err := protocol.NewSecureCodec(br, cli, secret)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if err := codec.WriteFrame(protocol.Frame{Type: protocol.FramePing, Payload: []byte("ENROLL " + sess.EnrollToken + " " + sess.ReconnectToken)}); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
f, err := codec.ReadFrame()
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if f.Type != protocol.FramePong || string(f.Payload) != "OK" {
|
|
t.Fatalf("unexpected ack: %#v", f)
|
|
}
|
|
}
|