diff --git a/ios/connect.go b/ios/connect.go index 27fb033f3..7ed1511b5 100755 --- a/ios/connect.go +++ b/ios/connect.go @@ -4,6 +4,7 @@ import ( "encoding/binary" "errors" "fmt" + "io" "net" "os" "strconv" @@ -184,15 +185,23 @@ func ConnectToServiceTunnelIface(device DeviceEntry, serviceName string) (Device } func CreateXpcConnection(h *http.HttpConnection) (*xpc.Connection, error) { - err := initializeXpcConnection(h) + clientServerChannel, err := h.OpenStream() + if err != nil { + return nil, fmt.Errorf("CreateXpcConnection: failed to open stream: %w", err) + } + serverClientChannel, err := h.OpenStream() + if err != nil { + return nil, fmt.Errorf("CreateXpcConnection: failed to open stream: %w", err) + } + err = initializeXpcConnection(clientServerChannel, serverClientChannel) if err != nil { return nil, fmt.Errorf("CreateXpcConnection: failed to initialize xpc connection: %w", err) } - clientServerChannel := http.NewStreamReadWriter(h, http.ClientServer) - serverClientChannel := http.NewStreamReadWriter(h, http.ServerClient) - - xpcConn, err := xpc.New(clientServerChannel, serverClientChannel, h) + openStream := func() (io.ReadWriteCloser, error) { + return h.OpenStream() + } + xpcConn, err := xpc.New(clientServerChannel, serverClientChannel, openStream, h) if err != nil { return nil, fmt.Errorf("CreateXpcConnection: failed to create xpc connection: %w", err) } @@ -247,10 +256,7 @@ func ConnectLockdownWithSession(device DeviceEntry) (*LockDownConnection, error) return lockdownConnection, nil } -func initializeXpcConnection(h *http.HttpConnection) error { - csWriter := http.NewStreamReadWriter(h, http.ClientServer) - ssWriter := http.NewStreamReadWriter(h, http.ServerClient) - +func initializeXpcConnection(csWriter, ssWriter io.ReadWriter) error { err := xpc.EncodeMessage(csWriter, xpc.Message{ Flags: xpc.AlwaysSetFlag, Body: map[string]interface{}{}, diff --git a/ios/http/http.go b/ios/http/http.go index 194add67c..0fbed0878 100644 --- a/ios/http/http.go +++ b/ios/http/http.go @@ -4,7 +4,7 @@ import ( "bytes" "fmt" "io" - "sync/atomic" + "sync" "github.com/danielpaulus/go-ios/ios/golog" "golang.org/x/net/http2" @@ -20,14 +20,61 @@ const ( ServerClient = StreamId(3) ) +// defaultWindowSize is the initial HTTP/2 flow-control window (RFC 9113 6.9.2) +// that applies until the peer announces a different one. +const defaultWindowSize = 65535 + +// recvWindowSize is the flow-control window we grant the peer, for the +// connection as well as for every stream. It is announced as +// SETTINGS_INITIAL_WINDOW_SIZE and replenished with WINDOW_UPDATE frames. +const recvWindowSize = 1048576 + +// windowUpdateThreshold is how much of a receive window may be used up before +// it is replenished, so that not every DATA frame needs a WINDOW_UPDATE. +const windowUpdateThreshold = recvWindowSize / 2 + +// maxFrameSize is the largest DATA frame payload we send. It's the HTTP/2 +// default for SETTINGS_MAX_FRAME_SIZE, which every peer has to accept. +const maxFrameSize = 16384 + // HttpConnection is a wrapper around a http2.Framer that provides a simple interface to read and write http2 streams for iOS17+. type HttpConnection struct { - framer *http2.Framer - clientServerStream *bytes.Buffer - serverClientStream *bytes.Buffer - closer io.Closer - csIsOpen *atomic.Bool - scIsOpen *atomic.Bool + closer io.Closer + + framer *http2.Framer + + // framerReadMu guards reading frames. A http2.Framer is not safe for concurrent + // reads, and the payload of a frame is only valid until the next read, so a + // frame has to be dispatched to its stream before the next one is read. + framerReadMu sync.Mutex + // framerWriteMu guards writing frames, a http2.Framer encodes every frame into + // the same buffer. + framerWriteMu sync.Mutex + + // mu guards the flow-control state and the additional streams below. + mu sync.Mutex + // peerInitialWindow is the peer's SETTINGS_INITIAL_WINDOW_SIZE, the send + // window every newly opened stream starts with. + peerInitialWindow int64 + // connSendWindow is the connection level send window. Only writes on + // additional streams wait for it, see Stream.Write. + connSendWindow int64 + // connRecvUnacked counts the bytes received on the connection that have not + // been granted back to the peer with a WINDOW_UPDATE yet. + connRecvUnacked int64 + // streams holds the additional client initiated streams (5, 7, ...) that + // are used for XPC file transfers. + streams map[uint32]*streamState + nextStreamId uint32 +} + +type streamState struct { + buf bytes.Buffer + sendWindow int64 + // recvUnacked counts the bytes received on this stream that have not been + // granted back to the peer with a WINDOW_UPDATE yet. + recvUnacked int64 + reset bool } func (r *HttpConnection) Close() error { @@ -44,13 +91,15 @@ func NewHttpConnection(rw io.ReadWriteCloser) (*HttpConnection, error) { err = framer.WriteSettings( http2.Setting{ID: http2.SettingMaxConcurrentStreams, Val: 100}, - http2.Setting{ID: http2.SettingInitialWindowSize, Val: 1048576}, + http2.Setting{ID: http2.SettingInitialWindowSize, Val: recvWindowSize}, ) if err != nil { return nil, fmt.Errorf("NewHttpConnection: could not write settings. %w", err) } - err = framer.WriteWindowUpdate(uint32(InitStream), 983041) + // SETTINGS_INITIAL_WINDOW_SIZE doesn't apply to the connection window, so it + // is raised to the same size explicitly (RFC 9113 6.9.2) + err = framer.WriteWindowUpdate(uint32(InitStream), recvWindowSize-defaultWindowSize) if err != nil { return nil, fmt.Errorf("NewHttpConnection: could not write window update. %w", err) } @@ -59,11 +108,13 @@ func NewHttpConnection(rw io.ReadWriteCloser) (*HttpConnection, error) { if err != nil { return nil, fmt.Errorf("NewHttpConnection: could not read frame. %w", err) } + peerInitialWindow := int64(defaultWindowSize) if frame.Header().Type == http2.FrameSettings { settings := frame.(*http2.SettingsFrame) v, ok := settings.Value(http2.SettingInitialWindowSize) if ok { framer.SetMaxReadFrameSize(v) + peerInitialWindow = int64(v) } err := framer.WriteSettingsAck() if err != nil { @@ -74,129 +125,268 @@ func NewHttpConnection(rw io.ReadWriteCloser) (*HttpConnection, error) { } return &HttpConnection{ - framer: framer, - clientServerStream: bytes.NewBuffer(nil), - serverClientStream: bytes.NewBuffer(nil), - closer: rw, - csIsOpen: &atomic.Bool{}, - scIsOpen: &atomic.Bool{}, + framer: framer, + closer: rw, + peerInitialWindow: peerInitialWindow, + connSendWindow: defaultWindowSize, + streams: map[uint32]*streamState{}, + nextStreamId: uint32(1), }, nil } -func (r *HttpConnection) ReadClientServerStream(p []byte) (int, error) { - for r.clientServerStream.Len() < len(p) { - err := r.readDataFrame() - if err != nil { - return 0, fmt.Errorf("ReadClientServerStream: %w", err) +// processFrame reads and handles a single frame. Callers waiting for stream +// data or for flow-control windows to open call it repeatedly until their +// condition is met. +// +// Only one goroutine reads from the connection at a time, and it dispatches the +// frames of all streams. ready is therefore evaluated again once this goroutine +// owns the read side: another goroutine may have delivered what the caller is +// waiting for in the meantime, and reading one more frame would block until the +// peer happens to send another one. +func (r *HttpConnection) processFrame(ready func() bool) error { + r.framerReadMu.Lock() + defer r.framerReadMu.Unlock() + if ready() { + return nil + } + f, err := r.framer.ReadFrame() + if err != nil { + return fmt.Errorf("could not read frame. %w", err) + } + switch f.Header().Type { + case http2.FrameData: + d := f.(*http2.DataFrame) + // the whole frame payload counts against the receive windows, the + // padding included (RFC 9113 6.9.1) + size := int64(d.Header().Length) + r.mu.Lock() + s, ok := r.streams[d.StreamID] + streamIncrement := int64(0) + if ok { + s.buf.Write(d.Data()) + // a stream that the peer is done with doesn't need its window back + if !s.reset && !d.StreamEnded() { + streamIncrement = ackReceived(&s.recvUnacked, size) + } + } + connIncrement := ackReceived(&r.connRecvUnacked, size) + r.mu.Unlock() + if !ok { + return fmt.Errorf("unknown stream id %d", d.StreamID) + } + // the received data is buffered without a bound, so the windows are + // replenished right away instead of when a reader consumes the data + if err := r.writeWindowUpdate(uint32(InitStream), connIncrement); err != nil { + return err + } + if err := r.writeWindowUpdate(d.StreamID, streamIncrement); err != nil { + return err + } + case http2.FrameGoAway: + return fmt.Errorf("received GOAWAY") + case http2.FrameSettings: + s := f.(*http2.SettingsFrame) + if s.Flags&http2.FlagSettingsAck != http2.FlagSettingsAck { + if v, ok := s.Value(http2.SettingInitialWindowSize); ok { + r.updateInitialWindow(int64(v)) + } + r.framerWriteMu.Lock() + err := r.framer.WriteSettingsAck() + r.framerWriteMu.Unlock() + if err != nil { + return fmt.Errorf("could not write settings ack. %w", err) + } + } + case http2.FrameWindowUpdate: + w := f.(*http2.WindowUpdateFrame) + r.mu.Lock() + if w.StreamID == uint32(InitStream) { + r.connSendWindow += int64(w.Increment) + } else if s, ok := r.streams[w.StreamID]; ok { + s.sendWindow += int64(w.Increment) + } + r.mu.Unlock() + case http2.FrameRSTStream: + rst := f.(*http2.RSTStreamFrame) + r.mu.Lock() + s, ok := r.streams[rst.StreamID] + if ok { + s.reset = true + } + r.mu.Unlock() + // The device resets file transfer streams once it received all data. + // That's no reason to fail reads on the XPC streams. + if !ok { + return fmt.Errorf("got RST frame with error code: %s", rst.ErrCode.String()) } + default: + break } - return r.clientServerStream.Read(p) + return nil } -func (r *HttpConnection) WriteClientServerStream(p []byte) (int, error) { - return r.write(p, uint32(ClientServer), r.csIsOpen) +// updateInitialWindow applies a new SETTINGS_INITIAL_WINDOW_SIZE of the peer, +// which also adjusts the send windows of all open streams (RFC 9113 6.9.2). +func (r *HttpConnection) updateInitialWindow(v int64) { + r.mu.Lock() + defer r.mu.Unlock() + delta := v - r.peerInitialWindow + r.peerInitialWindow = v + for _, s := range r.streams { + s.sendWindow += delta + } } -func (r *HttpConnection) WriteServerClientStream(p []byte) (int, error) { - return r.write(p, uint32(ServerClient), r.scIsOpen) +// ackReceived adds the n received bytes to unacked and returns the increment to +// grant back with a WINDOW_UPDATE frame, or 0 while the window still has room. +func ackReceived(unacked *int64, n int64) int64 { + *unacked += n + if *unacked < windowUpdateThreshold { + return 0 + } + increment := *unacked + *unacked = 0 + return increment } -func (r *HttpConnection) write(p []byte, stream uint32, isOpen *atomic.Bool) (int, error) { - if isOpen.CompareAndSwap(false, true) { - err := r.framer.WriteHeaders(http2.HeadersFrameParam{ - StreamID: stream, - EndHeaders: true, - }) - if err != nil { - return 0, fmt.Errorf("write: could not send headers. %w", err) - } +// writeWindowUpdate grants increment bytes of receive window back to the peer. +// The WINDOW_UPDATE is only sent when increment > 0, otherwise this is a no-op. +func (r *HttpConnection) writeWindowUpdate(streamId uint32, increment int64) error { + if increment <= 0 { + return nil } - return r.Write(p, stream) + r.framerWriteMu.Lock() + err := r.framer.WriteWindowUpdate(streamId, uint32(increment)) + r.framerWriteMu.Unlock() + if err != nil { + return fmt.Errorf("could not write window update for stream %d. %w", streamId, err) + } + return nil } -func (r *HttpConnection) Write(p []byte, streamId uint32) (int, error) { - err := r.framer.WriteData(streamId, false, p) +// Stream is an additional client initiated HTTP/2 stream. RemoteXPC uses those +// for transferring the payload of file transfer objects. +// +// Writes honor the peer's flow-control windows and read frames from the +// connection while they wait for window updates. Every stream of a connection +// can be read and written from its own goroutine. +type Stream struct { + h *HttpConnection + id uint32 +} + +// OpenStream opens a new client initiated stream. +func (r *HttpConnection) OpenStream() (*Stream, error) { + r.mu.Lock() + id := r.nextStreamId + // Client controlled streams are always odd numbered (RFC 9113 5.1.1) + r.nextStreamId += 2 + r.streams[id] = &streamState{sendWindow: r.peerInitialWindow} + r.mu.Unlock() + + r.framerWriteMu.Lock() + err := r.framer.WriteHeaders(http2.HeadersFrameParam{ + StreamID: id, + EndHeaders: true, + }) + r.framerWriteMu.Unlock() if err != nil { - return 0, fmt.Errorf("Write: could not write data. %w", err) + return nil, fmt.Errorf("OpenStream: could not send headers for stream %d. %w", id, err) } - return len(p), nil + return &Stream{h: r, id: id}, nil } -func (r *HttpConnection) readDataFrame() error { +// Read blocks until len(p) bytes were received on this stream +func (s *Stream) Read(p []byte) (int, error) { for { - f, err := r.framer.ReadFrame() - if err != nil { - return fmt.Errorf("readDataFrame: could not read frame. %w", err) - } - switch f.Header().Type { - case http2.FrameData: - d := f.(*http2.DataFrame) - switch d.StreamID { - case 1: - r.clientServerStream.Write(d.Data()) - case 3: - r.serverClientStream.Write(d.Data()) - default: - return fmt.Errorf("readDataFrame: unknown stream id %d", d.StreamID) - } - return nil - case http2.FrameGoAway: - return fmt.Errorf("received GOAWAY") - case http2.FrameSettings: - s := f.(*http2.SettingsFrame) - if s.Flags&http2.FlagSettingsAck != http2.FlagSettingsAck { - err := r.framer.WriteSettingsAck() - if err != nil { - return fmt.Errorf("readDataFrame: could not write settings ack. %w", err) - } - } - case http2.FrameRSTStream: - r := f.(*http2.RSTStreamFrame) - return fmt.Errorf("readDataFrame: got RST frame with error code: %s", r.ErrCode.String()) - default: - break + s.h.mu.Lock() + st := s.h.streams[s.id] + if st.buf.Len() >= len(p) { + n, err := st.buf.Read(p) + s.h.mu.Unlock() + return n, err + } + reset := st.reset + s.h.mu.Unlock() + if reset { + return 0, fmt.Errorf("Read: stream %d was reset by the peer", s.id) + } + if err := s.h.processFrame(func() bool { return s.readable(len(p)) }); err != nil { + return 0, fmt.Errorf("Read: %w", err) } } } -func (r *HttpConnection) ReadServerClientStream(p []byte) (int, error) { - for r.serverClientStream.Len() < len(p) { - err := r.readDataFrame() +// readable reports whether a read of n bytes can complete, either because the +// data arrived or because the peer reset the stream. +func (s *Stream) readable(n int) bool { + s.h.mu.Lock() + defer s.h.mu.Unlock() + st := s.h.streams[s.id] + return st.reset || st.buf.Len() >= n +} + +// Write sends p as DATA frames on this stream, waiting for the peer to open its +// flow-control windows whenever they are exhausted. +func (s *Stream) Write(p []byte) (int, error) { + written := 0 + for written < len(p) { + n, err := s.sendWindow(len(p) - written) + if err != nil { + return written, fmt.Errorf("Write: %w", err) + } + if n == 0 { + if err := s.h.processFrame(s.writable); err != nil { + return written, fmt.Errorf("Write: failed waiting for window update. %w", err) + } + continue + } + s.h.framerWriteMu.Lock() + err = s.h.framer.WriteData(s.id, false, p[written:written+n]) + s.h.framerWriteMu.Unlock() if err != nil { - return 0, err + return written, fmt.Errorf("Write: could not write data on stream %d. %w", s.id, err) } + written += n } - return r.serverClientStream.Read(p) -} - -type HttpStreamReadWriter struct { - h *HttpConnection - streamId uint32 + return written, nil } -func NewStreamReadWriter(h *HttpConnection, streamId StreamId) HttpStreamReadWriter { - return HttpStreamReadWriter{ - h: h, - streamId: uint32(streamId), - } +// writable reports whether a write can make progress, either because both send +// windows have room or because the peer reset the stream. +func (s *Stream) writable() bool { + s.h.mu.Lock() + defer s.h.mu.Unlock() + st := s.h.streams[s.id] + return st.reset || (st.sendWindow > 0 && s.h.connSendWindow > 0) } -func (h HttpStreamReadWriter) Read(p []byte) (n int, err error) { - if h.streamId == 1 { - return h.h.ReadClientServerStream(p) +// sendWindow reserves up to want bytes of the stream and connection windows and +// returns how many bytes may be sent right now. +func (s *Stream) sendWindow(want int) (int, error) { + s.h.mu.Lock() + defer s.h.mu.Unlock() + st := s.h.streams[s.id] + if st.reset { + return 0, fmt.Errorf("stream %d was reset by the peer", s.id) } - if h.streamId == 3 { - return h.h.ReadServerClientStream(p) + n := int64(min(want, maxFrameSize)) + n = min(n, st.sendWindow, s.h.connSendWindow) + if n <= 0 { + return 0, nil } - return 0, fmt.Errorf("Read: unknown stream id %d", h.streamId) + st.sendWindow -= n + s.h.connSendWindow -= n + return int(n), nil } -func (h HttpStreamReadWriter) Write(p []byte) (n int, err error) { - if h.streamId == 1 { - return h.h.WriteClientServerStream(p) - } - if h.streamId == 3 { - return h.h.WriteServerClientStream(p) +// Close half-closes the stream by sending an empty DATA frame with END_STREAM +func (s *Stream) Close() error { + s.h.framerWriteMu.Lock() + err := s.h.framer.WriteData(s.id, true, nil) + s.h.framerWriteMu.Unlock() + if err != nil { + return fmt.Errorf("Close: could not end stream %d. %w", s.id, err) } - return 0, fmt.Errorf("Write: unknown stream id %d", h.streamId) + return nil } diff --git a/ios/xpc/encoding.go b/ios/xpc/encoding.go index f7e2dadaf..71035eb94 100644 --- a/ios/xpc/encoding.go +++ b/ios/xpc/encoding.go @@ -44,6 +44,7 @@ const ( HeartbeatRequestFlag = uint32(0x00010000) HeartbeatReplyFlag = uint32(0x00020000) FileOpenFlag = uint32(0x00100000) + FileOpenReplyFlag = uint32(0x00200000) InitHandshakeFlag = uint32(0x00400000) ) @@ -156,6 +157,7 @@ func decodeWrapper(r io.Reader) (Message, error) { if h.BodyLen == 0 { return Message{ Flags: h.Flags, + Id: h.MsgId, }, nil } body, err := decodeBody(r, h) @@ -165,6 +167,7 @@ func decodeWrapper(r io.Reader) (Message, error) { return Message{ Flags: h.Flags, Body: body, + Id: h.MsgId, }, nil } @@ -551,6 +554,10 @@ func encodeObject(w io.Writer, e interface{}) error { if err := encodeDictionary(w, e.(map[string]interface{})); err != nil { return err } + case FileTransfer: + if err := encodeFileTransfer(w, t); err != nil { + return err + } default: return fmt.Errorf("can not encode type %v", t) } @@ -569,6 +576,24 @@ func encodeUuid(w io.Writer, u uuid.UUID) error { return nil } +// encodeFileTransfer writes a file transfer object. The payload itself is not +// part of the message, it gets sent on a separate stream that is opened with a +// FileOpenFlag message carrying the same MsgId. +func encodeFileTransfer(w io.Writer, f FileTransfer) error { + header := struct { + t xpcType + msgId uint64 + }{fileTransferType, f.MsgId} + if err := binary.Write(w, binary.LittleEndian, header); err != nil { + return fmt.Errorf("encodeFileTransfer: failed to write header: %w", err) + } + // the transfer length is always stored in a property 's' + if err := encodeDictionary(w, map[string]interface{}{"s": f.TransferSize}); err != nil { + return fmt.Errorf("encodeFileTransfer: failed to write transfer size: %w", err) + } + return nil +} + func encodeArray(w io.Writer, slice []interface{}) error { buf := bytes.NewBuffer(nil) for i, e := range slice { diff --git a/ios/xpc/encoding_test.go b/ios/xpc/encoding_test.go index a04d72f92..92296b882 100644 --- a/ios/xpc/encoding_test.go +++ b/ios/xpc/encoding_test.go @@ -3,6 +3,7 @@ package xpc import ( "bytes" "encoding/base64" + "encoding/hex" "github.com/google/uuid" "github.com/stretchr/testify/assert" "os" @@ -62,6 +63,7 @@ func TestDictionary(t *testing.T) { }, "CoreDevice.invocationIdentifier": "62419FC1-5ABF-4D96-BCA8-7A5F6F9A69EE", }, + Id: 1, }, res) } @@ -126,6 +128,13 @@ func TestEncodeDecode(t *testing.T) { }, expectedFlags: AlwaysSetFlag | DataFlag, }, + { + name: "encode file transfer", + input: map[string]interface{}{ + "image": FileTransfer{MsgId: 13, TransferSize: 16648704}, + }, + expectedFlags: AlwaysSetFlag | DataFlag, + }, { name: "encode uuid", input: map[string]interface{}{ @@ -164,3 +173,12 @@ func TestEncodeDecode(t *testing.T) { }) } } + +// TestEncodeFileTransferWireFormat checks the encoding against the bytes a Mac +// sends for the 'image' argument of the cryptexd 'install' routine. +func TestEncodeFileTransferWireFormat(t *testing.T) { + buf := bytes.NewBuffer(nil) + err := encodeObject(buf, FileTransfer{MsgId: 13, TransferSize: 16648704}) + assert.NoError(t, err) + assert.Equal(t, "00a001000d0000000000000000f0000014000000010000007300000000400000000afe0000000000", hex.EncodeToString(buf.Bytes())) +} diff --git a/ios/xpc/xpc.go b/ios/xpc/xpc.go index 3ff2cb150..9d0009ed9 100644 --- a/ios/xpc/xpc.go +++ b/ios/xpc/xpc.go @@ -5,26 +5,29 @@ package xpc import ( "fmt" "io" - - "golang.org/x/net/http2" ) // Connection represents a http2 based connection to an XPC service on an iOS17 device. type Connection struct { connectionCloser io.Closer - framer *http2.Framer msgId uint64 clientServer io.ReadWriter serverClient io.ReadWriter + openStream StreamOpener } +// StreamOpener opens an additional stream on the underlying connection. RemoteXPC +// sends the payload of each FileTransfer object on a stream of its own. +type StreamOpener func() (io.ReadWriteCloser, error) + // New creates a new connection to an XPC service on an iOS17 device. -func New(clientServer io.ReadWriter, serverClient io.ReadWriter, closer io.Closer) (*Connection, error) { +func New(clientServer io.ReadWriter, serverClient io.ReadWriter, openStream StreamOpener, closer io.Closer) (*Connection, error) { return &Connection{ connectionCloser: closer, msgId: 1, clientServer: clientServer, serverClient: serverClient, + openStream: openStream, }, nil } @@ -66,6 +69,55 @@ func (c *Connection) Send(data map[string]interface{}, flags ...uint32) error { return EncodeMessage(c.clientServer, msg) } +// FileTransferStream carries the payload of a FileTransfer object +type FileTransferStream struct { + rwc io.ReadWriteCloser + id uint64 +} + +// OpenFileTransfer opens a new stream for the payload of the FileTransfer object with +// the given id. The peer accepts it only after it received the message that +// references the FileTransfer, so call WaitAccepted after sending that message. +func (c *Connection) OpenFileTransfer(id uint64) (*FileTransferStream, error) { + if c.openStream == nil { + return nil, fmt.Errorf("OpenFileTransfer: connection does not support file transfers") + } + rwc, err := c.openStream() + if err != nil { + return nil, fmt.Errorf("OpenFileTransfer: %w", err) + } + err = EncodeMessage(rwc, Message{ + Flags: AlwaysSetFlag | FileOpenFlag, + Id: id, + }) + if err != nil { + return nil, fmt.Errorf("OpenFileTransfer: failed to send file open message: %w", err) + } + return &FileTransferStream{rwc: rwc, id: id}, nil +} + +// WaitAccepted blocks until the peer is ready to receive the payload +func (f *FileTransferStream) WaitAccepted() error { + msg, err := DecodeMessage(f.rwc) + if err != nil { + return fmt.Errorf("WaitAccepted: %w", err) + } + if msg.Flags&FileOpenReplyFlag == 0 || msg.Id != f.id { + return fmt.Errorf("WaitAccepted: unexpected reply with flags 0x%x and id %d for file transfer %d", msg.Flags, msg.Id, f.id) + } + return nil +} + +// Write sends payload data +func (f *FileTransferStream) Write(p []byte) (int, error) { + return f.rwc.Write(p) +} + +// Close marks the end of the payload +func (f *FileTransferStream) Close() error { + return f.rwc.Close() +} + func (c *Connection) Close() error { return c.connectionCloser.Close() }