@@ -689,6 +689,139 @@ func sessionIDFromParams(t *testing.T, params json.RawMessage) string {
689689 return decoded .SessionID
690690}
691691
692+ func TestClient_CreateSessionFailureClosesRegisteredSession (t * testing.T ) {
693+ tests := []struct {
694+ name string
695+ response func (string ) (json.RawMessage , * jsonrpc2.Error )
696+ wantErrSub string
697+ }{
698+ {
699+ name : "RPC failure" ,
700+ response : func (string ) (json.RawMessage , * jsonrpc2.Error ) {
701+ return nil , & jsonrpc2.Error {Code : - 32000 , Message : "session creation failed" }
702+ },
703+ wantErrSub : "failed to create session" ,
704+ },
705+ {
706+ name : "invalid response" ,
707+ response : func (string ) (json.RawMessage , * jsonrpc2.Error ) {
708+ return json .RawMessage (`"invalid"` ), nil
709+ },
710+ wantErrSub : "failed to unmarshal response" ,
711+ },
712+ {
713+ name : "session ID mismatch" ,
714+ response : func (string ) (json.RawMessage , * jsonrpc2.Error ) {
715+ return json .RawMessage (`{"sessionId":"different-session"}` ), nil
716+ },
717+ wantErrSub : "but the caller requested failed-session" ,
718+ },
719+ }
720+
721+ for _ , tt := range tests {
722+ t .Run (tt .name , func (t * testing.T ) {
723+ rpcClient , server , _ := newRuntimeShutdownRpcPair (t )
724+ t .Cleanup (server .Stop )
725+ client := & Client {
726+ client : rpcClient ,
727+ RPC : rpc .NewServerRPC (rpcClient ),
728+ sessions : make (map [string ]* Session ),
729+ }
730+
731+ captured := make (chan * Session , 1 )
732+ server .SetRequestHandler ("session.create" , func (params json.RawMessage ) (json.RawMessage , * jsonrpc2.Error ) {
733+ sessionID := sessionIDFromParams (t , params )
734+ client .sessionsMux .Lock ()
735+ session := client .sessions [sessionID ]
736+ client .sessionsMux .Unlock ()
737+ captured <- session
738+ return tt .response (sessionID )
739+ })
740+
741+ _ , err := client .CreateSession (t .Context (), & SessionConfig {SessionID : "failed-session" })
742+ if err == nil || ! strings .Contains (err .Error (), tt .wantErrSub ) {
743+ t .Fatalf ("CreateSession error = %v, want substring %q" , err , tt .wantErrSub )
744+ }
745+
746+ session := <- captured
747+ if session == nil {
748+ t .Fatal ("session was not registered before session.create" )
749+ }
750+ assertSessionEventChannelClosed (t , session )
751+ assertSessionNotRegistered (t , client , "failed-session" )
752+ })
753+ }
754+ }
755+
756+ func TestClient_CreateSessionInitializationFailureClosesRegisteredSession (t * testing.T ) {
757+ rpcClient , server , _ := newRuntimeShutdownRpcPair (t )
758+ t .Cleanup (server .Stop )
759+ client := & Client {
760+ client : rpcClient ,
761+ RPC : rpc .NewServerRPC (rpcClient ),
762+ sessions : make (map [string ]* Session ),
763+ options : ClientOptions {SessionFS : & SessionFSConfig {
764+ InitialWorkingDirectory : "/" ,
765+ SessionStatePath : "/session-state" ,
766+ Conventions : rpc .SessionFSSetProviderConventionsPosix ,
767+ Capabilities : & SessionFSCapabilities {Sqlite : true },
768+ }},
769+ }
770+
771+ var captured * Session
772+ _ , err := client .CreateSession (t .Context (), & SessionConfig {
773+ SessionID : "failed-session-fs" ,
774+ CreateSessionFSProvider : func (session * Session ) SessionFSProvider {
775+ captured = session
776+ return noSQLiteSessionFSProvider {}
777+ },
778+ })
779+ if err == nil || ! strings .Contains (err .Error (), "does not implement SessionFSSqliteProvider" ) {
780+ t .Fatalf ("CreateSession error = %v, want SQLite provider validation error" , err )
781+ }
782+ if captured == nil {
783+ t .Fatal ("CreateSessionFSProvider did not receive the registered session" )
784+ }
785+ assertSessionEventChannelClosed (t , captured )
786+ assertSessionNotRegistered (t , client , "failed-session-fs" )
787+ }
788+
789+ func assertSessionEventChannelClosed (t * testing.T , session * Session ) {
790+ t .Helper ()
791+ select {
792+ case _ , ok := <- session .eventCh :
793+ if ok {
794+ t .Fatal ("session event channel is still open" )
795+ }
796+ case <- time .After (time .Second ):
797+ t .Fatal ("timed out waiting for session event channel to close" )
798+ }
799+ }
800+
801+ func assertSessionNotRegistered (t * testing.T , client * Client , sessionID string ) {
802+ t .Helper ()
803+ client .sessionsMux .Lock ()
804+ defer client .sessionsMux .Unlock ()
805+ if _ , ok := client .sessions [sessionID ]; ok {
806+ t .Fatalf ("session %q is still registered" , sessionID )
807+ }
808+ }
809+
810+ type noSQLiteSessionFSProvider struct {}
811+
812+ func (noSQLiteSessionFSProvider ) ReadFile (string ) (string , error ) { return "" , nil }
813+ func (noSQLiteSessionFSProvider ) WriteFile (string , string , * int ) error { return nil }
814+ func (noSQLiteSessionFSProvider ) AppendFile (string , string , * int ) error { return nil }
815+ func (noSQLiteSessionFSProvider ) Exists (string ) (bool , error ) { return false , nil }
816+ func (noSQLiteSessionFSProvider ) Stat (string ) (* SessionFSFileInfo , error ) { return nil , nil }
817+ func (noSQLiteSessionFSProvider ) MakeDirectory (string , bool , * int ) error { return nil }
818+ func (noSQLiteSessionFSProvider ) ReadDirectory (string ) ([]string , error ) { return nil , nil }
819+ func (noSQLiteSessionFSProvider ) ReadDirectoryWithTypes (string ) ([]rpc.SessionFSReaddirWithTypesEntry , error ) {
820+ return nil , nil
821+ }
822+ func (noSQLiteSessionFSProvider ) Remove (string , bool , bool ) error { return nil }
823+ func (noSQLiteSessionFSProvider ) Rename (string , string ) error { return nil }
824+
692825func assertRuntimeShutdownNotCalled (t * testing.T , shutdownCalled <- chan struct {}) {
693826 t .Helper ()
694827 select {
0 commit comments