mirror of
https://github.com/gravitational/teleport.git
synced 2026-09-24 16:17:11 +08:00
Fix invalid Write implementation on K8S join stream (#15503)
This commit is contained in:
@@ -215,12 +215,11 @@ func (s *SessionStream) Write(data []byte) (int, error) {
|
||||
defer s.writeSync.Unlock()
|
||||
|
||||
err := s.conn.WriteMessage(websocket.BinaryMessage, data)
|
||||
|
||||
if err != nil {
|
||||
return len(data), s.conn.WriteMessage(websocket.BinaryMessage, data)
|
||||
return 0, trace.Wrap(err)
|
||||
}
|
||||
|
||||
return 0, trace.Wrap(err)
|
||||
return len(data), nil
|
||||
}
|
||||
|
||||
// Resize sends a resize request to the other party.
|
||||
|
||||
@@ -0,0 +1,113 @@
|
||||
/*
|
||||
Copyright 2022 Gravitational, Inc.
|
||||
|
||||
Licensed under the Apache License, Version 2.0 (the "License");
|
||||
you may not use this file except in compliance with the License.
|
||||
You may obtain a copy of the License at
|
||||
|
||||
http://www.apache.org/licenses/LICENSE-2.0
|
||||
|
||||
Unless required by applicable law or agreed to in writing, software
|
||||
distributed under the License is distributed on an "AS IS" BASIS,
|
||||
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
See the License for the specific language governing permissions and
|
||||
limitations under the License.
|
||||
*/
|
||||
|
||||
package streamproto
|
||||
|
||||
import (
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/gorilla/websocket"
|
||||
"github.com/gravitational/teleport/api/types"
|
||||
"github.com/gravitational/trace"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
var upgrader = websocket.Upgrader{}
|
||||
|
||||
func TestPingPong(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
runClient := func(conn *websocket.Conn) error {
|
||||
client, err := NewSessionStream(conn, ClientHandshake{Mode: types.SessionPeerMode})
|
||||
if err != nil {
|
||||
return trace.Wrap(err)
|
||||
}
|
||||
|
||||
n, err := client.Write([]byte("ping"))
|
||||
if err != nil {
|
||||
return trace.Wrap(err)
|
||||
}
|
||||
if n != 4 {
|
||||
return trace.Errorf("unexpected write size: %d", n)
|
||||
}
|
||||
|
||||
out := make([]byte, 4)
|
||||
_, err = io.ReadFull(client, out)
|
||||
if err != nil {
|
||||
return trace.Wrap(err)
|
||||
}
|
||||
if string(out) != "pong" {
|
||||
return trace.BadParameter("expected pong, got %q", out)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
runServer := func(conn *websocket.Conn) error {
|
||||
server, err := NewSessionStream(conn, ServerHandshake{MFARequired: false})
|
||||
if err != nil {
|
||||
return trace.Wrap(err)
|
||||
}
|
||||
|
||||
out := make([]byte, 4)
|
||||
_, err = io.ReadFull(server, out)
|
||||
if err != nil {
|
||||
return trace.Wrap(err)
|
||||
}
|
||||
if string(out) != "ping" {
|
||||
return trace.BadParameter("expected ping, got %q", out)
|
||||
}
|
||||
|
||||
n, err := server.Write([]byte("pong"))
|
||||
if err != nil {
|
||||
return trace.Wrap(err)
|
||||
}
|
||||
if n != 4 {
|
||||
return trace.Errorf("unexpected write size: %d", n)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
errCh := make(chan error, 2)
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
ws, err := upgrader.Upgrade(w, r, nil)
|
||||
if err != nil {
|
||||
http.Error(w, err.Error(), http.StatusInternalServerError)
|
||||
}
|
||||
|
||||
defer ws.WriteControl(websocket.CloseMessage, websocket.FormatCloseMessage(websocket.CloseNormalClosure, ""), time.Time{})
|
||||
errCh <- runServer(ws)
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
url := "ws" + strings.TrimPrefix(server.URL, "http")
|
||||
ws, _, err := websocket.DefaultDialer.Dial(url, nil)
|
||||
require.NoError(t, err)
|
||||
|
||||
go func() {
|
||||
defer ws.Close()
|
||||
errCh <- runClient(ws)
|
||||
}()
|
||||
|
||||
require.NoError(t, <-errCh)
|
||||
require.NoError(t, <-errCh)
|
||||
}
|
||||
Reference in New Issue
Block a user