Add protocol v2 packet primitives · Entire
Add protocol v2 packet primitives
Sessions
e87357af16c9View transcript
Changes
2
internal/syncer
Aprotocol_v2.go+173
Aprotocol_v2_test.go+87
package syncer
import (
"bufio"
"bytes"
"fmt"
"io"
"strings"
"github.com/go-git/go-git/v5/plumbing/format/pktline"
)
const (
delimPkt = "0001"
responseEndPkt = "0002"
)
type packetType int
const (
packetTypeData packetType = iota
packetTypeFlush
packetTypeDelim
packetTypeResponseEnd
)
type v2CapabilityAdvertisement struct {
Capabilities map[string]string
}
func (a *v2CapabilityAdvertisement) Supports(name string) bool {
if a == nil {
return false
}
_, ok := a.Capabilities[name]
return ok
}
func (a *v2CapabilityAdvertisement) Value(name string) string {
if a == nil {
return ""
}
return a.Capabilities[name]
}
type packetReader struct {
r *bufio.Reader
}
func newPacketReader(r io.Reader) *packetReader {
if br, ok := r.(*bufio.Reader); ok {
return &packetReader{r: br}
}
return &packetReader{r: bufio.NewReader(r)}
}
func (r *packetReader) Reader() *bufio.Reader {
return r.r
}
func (r *packetReader) ReadPacket() (packetType, []byte, error) {
header := make([]byte, 4)
if _, err := io.ReadFull(r.r, header); err != nil {
return packetTypeData, nil, err
}
switch string(header) {
case "0000":
return packetTypeFlush, nil, nil
case delimPkt:
return packetTypeDelim, nil, nil
case responseEndPkt:
return packetTypeResponseEnd, nil, nil
}
var headerArr [4]byte
copy(headerArr[:], header)
n, err := pktlineLength(headerArr)
if err != nil {
return packetTypeData, nil, err
}
if n <= 4 {
return packetTypeData, nil, pktline.ErrInvalidPktLen
}
payload := make([]byte, n-4)
if _, err := io.ReadFull(r.r, payload); err != nil {
return packetTypeData, nil, err
}
return packetTypeData, payload, nil
}
func pktlineLength(header [4]byte) (int, error) {
var n int
for _, b := range header {
value, err := asciiHexToByte(b)
if err != nil {
return 0, pktline.ErrInvalidPktLen
}
n = 16*n + int(value)
}
return n, nil
}
func asciiHexToByte(b byte) (byte, error) {
switch {
case b >= '0' && b <= '9':
return b - '0', nil
case b >= 'a' && b <= 'f':
return b - 'a' + 10, nil
case b >= 'A' && b <= 'F':
return b - 'A' + 10, nil
default:
return 0, pktline.ErrInvalidPktLen
}
}
func decodeV2CapabilityAdvertisement(r io.Reader) (*v2CapabilityAdvertisement, error) {
reader := newPacketReader(r)
kind, payload, err := reader.ReadPacket()
if err != nil {
return nil, err
}
if kind != packetTypeData || string(payload) != "version 2\n" {
return nil, fmt.Errorf("unexpected protocol advertisement %q", payload)
}
adv := &v2CapabilityAdvertisement{Capabilities: map[string]string{}}
for {
kind, payload, err = reader.ReadPacket()
if err != nil {
return nil, err
}
if kind == packetTypeFlush {
return adv, nil
}
if kind != packetTypeData {
return nil, fmt.Errorf("unexpected packet type %v in capability advertisement", kind)
}
line := strings.TrimSuffix(string(payload), "\n")
name, value, _ := strings.Cut(line, "=")
adv.Capabilities[name] = value
}
}
func encodeV2CommandRequest(command string, capabilityArgs []string, commandArgs []string) ([]byte, error) {
var buf bytes.Buffer
enc := pktline.NewEncoder(&buf)
if err := enc.EncodeString("command=" + command + "\n"); err != nil {
return nil, err
}
for _, arg := range capabilityArgs {
if err := enc.EncodeString(arg + "\n"); err != nil {
return nil, err
}
}
if len(commandArgs) > 0 {
if _, err := buf.WriteString(delimPkt); err != nil {
return nil, err
}
for _, arg := range commandArgs {
if err := enc.EncodeString(arg + "\n"); err != nil {
return nil, err
}
}
}
if err := enc.Flush(); err != nil {
return nil, err
}
return buf.Bytes(), nil
}
Ainternal/syncer/protocol_v2.go+173
package syncer
import (
"bytes"
"testing"
)
func TestPacketReaderHandlesSpecialPackets(t *testing.T) {
reader := newPacketReader(bytes.NewBufferString("0000000100020006a\n"))
kind, payload, err := reader.ReadPacket()
if err != nil {
t.Fatalf("read flush: %v", err)
}
if kind != packetTypeFlush || payload != nil {
t.Fatalf("unexpected flush packet: kind=%v payload=%q", kind, payload)
}
kind, payload, err = reader.ReadPacket()
if err != nil {
t.Fatalf("read delim: %v", err)
}
if kind != packetTypeDelim {
t.Fatalf("unexpected delim kind: %v", kind)
}
kind, payload, err = reader.ReadPacket()
if err != nil {
t.Fatalf("read response-end: %v", err)
}
if kind != packetTypeResponseEnd {
t.Fatalf("unexpected response-end kind: %v", kind)
}
kind, payload, err = reader.ReadPacket()
if err != nil {
t.Fatalf("read data: %v", err)
}
if kind != packetTypeData || string(payload) != "a\n" {
t.Fatalf("unexpected data packet: kind=%v payload=%q", kind, payload)
}
}
func TestDecodeV2CapabilityAdvertisement(t *testing.T) {
wire := "" +
"000eversion 2\n" +
"0013ls-refs=unborn\n" +
"0012fetch=shallow\n" +
"0013agent=git/test\n" +
"0000"
adv, err := decodeV2CapabilityAdvertisement(bytes.NewBufferString(wire))
if err != nil {
t.Fatalf("decode advertisement: %v", err)
}
if !adv.Supports("ls-refs") {
t.Fatalf("expected ls-refs capability")
}
if got := adv.Value("fetch"); got != "shallow" {
t.Fatalf("unexpected fetch value %q", got)
}
if got := adv.Value("agent"); got != "git/test" {
t.Fatalf("unexpected agent value %q", got)
}
}
func TestEncodeV2CommandRequest(t *testing.T) {
req, err := encodeV2CommandRequest(
"ls-refs",
[]string{"agent=git-sync/test"},
[]string{"peel", "ref-prefix refs/heads/"},
)
if err != nil {
t.Fatalf("encode request: %v", err)
}
want := "" +
"0014command=ls-refs\n" +
"0018agent=git-sync/test\n" +
"0001" +
"0009peel\n" +
"001bref-prefix refs/heads/\n" +
"0000"
if string(req) != want {
t.Fatalf("unexpected request:\n%s\nwant:\n%s", req, want)
}
}
Ainternal/syncer/protocol_v2_test.go+87