@@ -944,3 +944,181 @@ multiDone:
944944 t .Log ("✅ Multi-tool pipeline verified: session → tokens → tool_call → tool_result → (repeat) → done" )
945945}
946946
947+ // TestServe_E2E_TokenStats verifies that the done event includes
948+ // latency, contextTokens, outputTokens, and session-level token
949+ // economics fields.
950+ func TestServe_E2E_TokenStats (t * testing.T ) {
951+ // Mock LLM server with usage info in every response
952+ callCount := 0
953+ llmSrv := httptest .NewServer (http .HandlerFunc (func (w http.ResponseWriter , r * http.Request ) {
954+ callCount ++
955+ w .Header ().Set ("Content-Type" , "application/json" )
956+ if callCount <= 1 {
957+ // First call: tool call with usage
958+ fmt .Fprint (w , `{"choices":[{"message":{"content":"Checking.","tool_calls":[{"id":"c_1","function":{"name":"shell","arguments":"{\"command\":\"echo ok\"}"}}]}}],"usage":{"prompt_tokens":200,"completion_tokens":30}}` )
959+ } else {
960+ // Final answer with usage
961+ fmt .Fprint (w , `{"choices":[{"message":{"content":"All done."}}],"usage":{"prompt_tokens":300,"completion_tokens":60}}` )
962+ }
963+ }))
964+ defer llmSrv .Close ()
965+
966+ envCleanup := setTestEnv (t , llmSrv .URL )
967+ defer envCleanup ()
968+
969+ store := newTestSessionStore (t )
970+ ln , mux := buildServeMux (t , store )
971+ defer ln .Close ()
972+
973+ errCh := make (chan error , 1 )
974+ go func () { errCh <- serveOnListener (ln , mux ) }()
975+ waitForHTTP (t , ln .Addr ().String ())
976+
977+ wsURL := "ws://" + ln .Addr ().String () + "/ws"
978+ conn , err := golangws .Dial (wsURL , "" , "http://localhost" )
979+ if err != nil {
980+ t .Fatalf ("Dial(%q): %v" , wsURL , err )
981+ }
982+ defer conn .Close ()
983+ t .Log ("✅ WebSocket connected" )
984+
985+ // Send first prompt
986+ prompt := map [string ]string {"type" : "prompt" , "content" : "run a command" }
987+ payload , _ := json .Marshal (prompt )
988+ if err := golangws .Message .Send (conn , string (payload )); err != nil {
989+ t .Fatalf ("Send: %v" , err )
990+ }
991+ t .Log ("✅ Prompt 1 sent" )
992+
993+ // Collect events — focus on the done event
994+ conn .SetReadDeadline (time .Now ().Add (15 * time .Second ))
995+
996+ var doneEvent map [string ]any
997+ for i := 0 ; i < 15 ; i ++ {
998+ var raw []byte
999+ if err := golangws .Message .Receive (conn , & raw ); err != nil {
1000+ t .Fatalf ("Receive event %d: %v" , i , err )
1001+ }
1002+ t .Logf (" event[%d]: %s" , i , string (raw ))
1003+
1004+ var evt map [string ]any
1005+ if err := json .Unmarshal (raw , & evt ); err != nil {
1006+ t .Fatalf ("unmarshal event %d: %v" , i , err )
1007+ }
1008+ if evt ["type" ] == "done" {
1009+ doneEvent = evt
1010+ goto statsCheck
1011+ }
1012+ if evt ["type" ] == "error" {
1013+ t .Fatalf ("unexpected error: %v" , evt ["message" ])
1014+ }
1015+ }
1016+ t .Fatal ("did not receive done event" )
1017+
1018+ statsCheck:
1019+ if doneEvent == nil {
1020+ t .Fatal ("done event not found" )
1021+ }
1022+
1023+ // Validate done event fields
1024+ latency , ok := doneEvent ["latency" ].(float64 )
1025+ if ! ok {
1026+ t .Error ("done event missing 'latency' field" )
1027+ } else if latency <= 0 {
1028+ t .Errorf ("latency = %v, want > 0" , latency )
1029+ }
1030+ t .Logf (" latency: %.2fs" , latency )
1031+
1032+ ctxTokens , ok := doneEvent ["contextTokens" ].(float64 )
1033+ if ! ok {
1034+ t .Error ("done event missing 'contextTokens' field" )
1035+ } else if ctxTokens != 500 { // 200 + 300
1036+ t .Errorf ("contextTokens = %.0f, want 500" , ctxTokens )
1037+ }
1038+ t .Logf (" contextTokens: %.0f" , ctxTokens )
1039+
1040+ outTokens , ok := doneEvent ["outputTokens" ].(float64 )
1041+ if ! ok {
1042+ t .Error ("done event missing 'outputTokens' field" )
1043+ } else if outTokens != 90 { // 30 + 60
1044+ t .Errorf ("outputTokens = %.0f, want 90" , outTokens )
1045+ }
1046+ t .Logf (" outputTokens: %.0f" , outTokens )
1047+
1048+ // Session-level stats (first prompt = same as turn-level)
1049+ sessCtx , ok := doneEvent ["sessionContextTokens" ].(float64 )
1050+ if ! ok {
1051+ t .Error ("done event missing 'sessionContextTokens' field" )
1052+ } else if sessCtx != 500 {
1053+ t .Errorf ("sessionContextTokens = %.0f, want 500" , sessCtx )
1054+ }
1055+ t .Logf (" sessionContextTokens: %.0f" , sessCtx )
1056+
1057+ sessOut , ok := doneEvent ["sessionOutputTokens" ].(float64 )
1058+ if ! ok {
1059+ t .Error ("done event missing 'sessionOutputTokens' field" )
1060+ } else if sessOut != 90 {
1061+ t .Errorf ("sessionOutputTokens = %.0f, want 90" , sessOut )
1062+ }
1063+ t .Logf (" sessionOutputTokens: %.0f" , sessOut )
1064+
1065+ // Send a second prompt to verify session-level accumulation
1066+ callCount = 0 // reset mock so next prompt also makes a tool call
1067+ prompt2 := map [string ]string {"type" : "prompt" , "content" : "do another thing" }
1068+ payload2 , _ := json .Marshal (prompt2 )
1069+ if err := golangws .Message .Send (conn , string (payload2 )); err != nil {
1070+ t .Fatalf ("Send prompt 2: %v" , err )
1071+ }
1072+ t .Log ("✅ Prompt 2 sent" )
1073+
1074+ conn .SetReadDeadline (time .Now ().Add (15 * time .Second ))
1075+ var done2 map [string ]any
1076+ for i := 0 ; i < 15 ; i ++ {
1077+ var raw []byte
1078+ if err := golangws .Message .Receive (conn , & raw ); err != nil {
1079+ t .Fatalf ("Receive event %d (prompt 2): %v" , i , err )
1080+ }
1081+ t .Logf (" event[%d]: %s" , i , string (raw ))
1082+
1083+ var evt map [string ]any
1084+ if err := json .Unmarshal (raw , & evt ); err != nil {
1085+ t .Fatalf ("unmarshal event %d: %v" , i , err )
1086+ }
1087+ if evt ["type" ] == "done" {
1088+ done2 = evt
1089+ goto sessionCheck
1090+ }
1091+ if evt ["type" ] == "error" {
1092+ t .Fatalf ("unexpected error: %v" , evt ["message" ])
1093+ }
1094+ }
1095+ t .Fatal ("did not receive second done event" )
1096+
1097+ sessionCheck:
1098+ if done2 == nil {
1099+ t .Fatal ("second done event not found" )
1100+ }
1101+
1102+ // Turn-level stats should reflect only this turn's tokens
1103+ ctx2 , _ := done2 ["contextTokens" ].(float64 )
1104+ if ctx2 != 500 {
1105+ t .Errorf ("prompt 2 contextTokens = %.0f, want 500" , ctx2 )
1106+ }
1107+ out2 , _ := done2 ["outputTokens" ].(float64 )
1108+ if out2 != 90 {
1109+ t .Errorf ("prompt 2 outputTokens = %.0f, want 90" , out2 )
1110+ }
1111+
1112+ // Session-level stats should be the sum of both turns
1113+ sessCtx2 , _ := done2 ["sessionContextTokens" ].(float64 )
1114+ if sessCtx2 != 1000 { // 500 + 500
1115+ t .Errorf ("sessionContextTokens (prompt 2) = %.0f, want 1000" , sessCtx2 )
1116+ }
1117+ sessOut2 , _ := done2 ["sessionOutputTokens" ].(float64 )
1118+ if sessOut2 != 180 { // 90 + 90
1119+ t .Errorf ("sessionOutputTokens (prompt 2) = %.0f, want 180" , sessOut2 )
1120+ }
1121+
1122+ t .Log ("✅ Token stats verified: turn-level + session-level accumulation" )
1123+ }
1124+
0 commit comments