per response decode to not bulk bail on bad entries

This commit is contained in:
its-a-feature
2026-05-11 20:49:33 -05:00
parent d97bc96e46
commit ff6e856988
4 changed files with 528 additions and 73 deletions
@@ -389,7 +389,7 @@ func processAgentMessageContent(agentMessageInput *AgentMessageRawInput, uuidInf
{
response, err = handleAgentMessageGetTasking(&decryptedMessage, uuidInfo.CallbackID)
instanceResponse.OuterUuid = uuidInfo.UUID // this is what our message UUID was coming into this parsing
if getDelegateTasks, ok := decryptedMessage["get_delegate_tasks"]; !ok || getDelegateTasks.(bool) {
if shouldAgentMessageGetDelegateTasks(decryptedMessage) {
// this means we should try to get some delegated tasks if they exist for our callback
delegateResponses = append(delegateResponses, getDelegateTaskMessages(uuidInfo.CallbackID, instanceResponse.AgentUUIDSize, agentMessageInput.UpdateCheckinTime)...)
} else {
@@ -1244,6 +1244,23 @@ func reflectBackOtherKeys(response *map[string]interface{}, other *map[string]in
}
}
// collectOtherKeys mirrors mapstructure's ",remain" behavior for handlers that
// decode only the fields they actually consume.
func collectOtherKeys(message map[string]interface{}, consumedKeys ...string) map[string]interface{} {
consumed := make(map[string]struct{}, len(consumedKeys))
for _, key := range consumedKeys {
consumed[key] = struct{}{}
}
other := make(map[string]interface{}, len(message))
for key, val := range message {
if _, ok := consumed[key]; ok {
continue
}
other[key] = val
}
return other
}
func GetUUIDBytes(outerUUID string, agentUUIDLength int) ([]byte, error) {
switch agentUUIDLength {
case 36:
@@ -2,15 +2,16 @@ package rabbitmq
import (
"database/sql"
"encoding/json"
"errors"
"fmt"
"math"
"time"
"github.com/its-a-feature/Mythic/database"
databaseStructs "github.com/its-a-feature/Mythic/database/structs"
"github.com/its-a-feature/Mythic/logging"
"github.com/jmoiron/sqlx"
"github.com/mitchellh/mapstructure"
)
type agentMessageGetTasking struct {
@@ -101,6 +102,100 @@ func buildAgentMessageTask(task agentMessageTaskRow) agentMessageGetTaskingTask
return newTask
}
// decodeAgentMessageGetTasking extracts only the Mythic-owned get_tasking fields
// while keeping agent-defined tracking keys available for the response.
func decodeAgentMessageGetTasking(incoming map[string]interface{}) (agentMessageGetTasking, error) {
agentMessage := agentMessageGetTasking{
GetDelegateTasks: true,
Other: collectOtherKeys(incoming, "tasking_size", "get_delegate_tasks"),
}
if taskingSize, ok := incoming["tasking_size"]; ok {
parsedTaskingSize, err := parseAgentMessageInt(taskingSize, "tasking_size")
if err != nil {
return agentMessage, err
}
agentMessage.TaskingSize = parsedTaskingSize
}
if getDelegateTasks, ok := incoming["get_delegate_tasks"]; ok {
parsedGetDelegateTasks, ok := getDelegateTasks.(bool)
if !ok {
return agentMessage, fmt.Errorf("get_delegate_tasks must be a bool, got %T", getDelegateTasks)
}
agentMessage.GetDelegateTasks = parsedGetDelegateTasks
}
return agentMessage, nil
}
func parseAgentMessageInt(value interface{}, field string) (int, error) {
switch typedValue := value.(type) {
case int:
return typedValue, nil
case int8:
return int(typedValue), nil
case int16:
return int(typedValue), nil
case int32:
return int(typedValue), nil
case int64:
if typedValue > int64(math.MaxInt) || typedValue < int64(math.MinInt) {
return 0, fmt.Errorf("%s is outside int range", field)
}
return int(typedValue), nil
case uint:
if uint64(typedValue) > uint64(math.MaxInt) {
return 0, fmt.Errorf("%s is outside int range", field)
}
return int(typedValue), nil
case uint8:
return int(typedValue), nil
case uint16:
return int(typedValue), nil
case uint32:
if uint64(typedValue) > uint64(math.MaxInt) {
return 0, fmt.Errorf("%s is outside int range", field)
}
return int(typedValue), nil
case uint64:
if typedValue > uint64(math.MaxInt) {
return 0, fmt.Errorf("%s is outside int range", field)
}
return int(typedValue), nil
case float64:
if math.Trunc(typedValue) != typedValue {
return 0, fmt.Errorf("%s must be an integer, got %v", field, typedValue)
}
if typedValue > float64(math.MaxInt) || typedValue < float64(math.MinInt) {
return 0, fmt.Errorf("%s is outside int range", field)
}
return int(typedValue), nil
case float32:
return parseAgentMessageInt(float64(typedValue), field)
case json.Number:
parsedValue, err := typedValue.Int64()
if err != nil {
return 0, fmt.Errorf("%s must be an integer: %w", field, err)
}
return parseAgentMessageInt(parsedValue, field)
default:
return 0, fmt.Errorf("%s must be an integer, got %T", field, value)
}
}
// shouldAgentMessageGetDelegateTasks keeps delegate tasking default-on while
// avoiding a panic if an agent sends a malformed get_delegate_tasks value.
func shouldAgentMessageGetDelegateTasks(incoming map[string]interface{}) bool {
getDelegateTasks, ok := incoming["get_delegate_tasks"]
if !ok {
return true
}
parsedGetDelegateTasks, ok := getDelegateTasks.(bool)
if !ok {
logging.LogError(nil, "Invalid get_delegate_tasks value in agent get_tasking message", "value", getDelegateTasks)
return true
}
return parsedGetDelegateTasks
}
func handleAgentMessageGetTasking(incoming *map[string]interface{}, callbackID int) (map[string]interface{}, error) {
// got message:
/*
@@ -116,7 +211,6 @@ func handleAgentMessageGetTasking(incoming *map[string]interface{}, callbackID i
1. check for direct tasks
2. check for delegate tasks
*/
agentMessage := agentMessageGetTasking{}
currentTasks := []databaseStructs.Task{}
if taskIDs := submittedTasksAwaitingFetching.getTasksForCallbackId(callbackID); len(taskIDs) > 0 {
query, args, err := sqlx.Named(`SELECT
@@ -141,7 +235,7 @@ func handleAgentMessageGetTasking(incoming *map[string]interface{}, callbackID i
}
}
err := mapstructure.Decode(incoming, &agentMessage)
agentMessage, err := decodeAgentMessageGetTasking(*incoming)
if err != nil {
logging.LogError(err, "Failed to decode agent message into struct")
return nil, errors.New(fmt.Sprintf("Failed to decode agent message into agentMessageGetTasking struct: %s", err.Error()))
@@ -60,9 +60,17 @@ const selectAgentMessagePostResponseTasksQuery = `SELECT
JOIN payload ON callback.registered_payload_id = payload.id
WHERE task.agent_task_id IN (?)`
type agentMessagePostResponseMessage struct {
Responses []agentMessagePostResponse `json:"responses" mapstructure:"responses" xml:"responses"`
Other map[string]interface{} `json:"-" mapstructure:",remain"` // capture any 'other' keys that were passed in so we can reply back with them
type decodedAgentMessagePostResponse struct {
Index int
Response agentMessagePostResponse
DecodeError error
TaskID string
Other map[string]interface{}
}
type decodedAgentMessagePostResponseMessage struct {
Responses []decodedAgentMessagePostResponse
Other map[string]interface{}
}
type agentMessagePostResponse struct {
@@ -93,6 +101,33 @@ type agentMessagePostResponse struct {
Other map[string]interface{} `json:"-" mapstructure:",remain"` // capture any 'other' keys that were passed in so we can reply back with them
}
var agentMessagePostResponseConsumedKeys = map[string]struct{}{
"alerts": {},
"artifacts": {},
"callback": {},
"callback_tokens": {},
"commands": {},
"completed": {},
"credentials": {},
"custom_browser": {},
"download": {},
"edges": {},
"events": {},
"file_browser": {},
"keylogs": {},
"process_response": {},
"processes": {},
"removed_files": {},
"sequence_num": {},
"status": {},
"stderr": {},
"stdout": {},
"task_id": {},
"tokens": {},
"upload": {},
"user_output": {},
}
var ValidCredentialTypesList = []string{"plaintext", "certificate", "hash", "key", "ticket", "cookie", "hex"}
type agentMessagePostResponseFileBrowser struct {
@@ -451,6 +486,137 @@ func getAgentMessagePostResponseTasks(responses []agentMessagePostResponse) (map
return tasksByAgentTaskID, nil
}
// decodeAgentMessagePostResponseMessage decodes each response entry separately
// so one malformed response cannot discard valid sibling responses.
func decodeAgentMessagePostResponseMessage(incoming map[string]interface{}) (decodedAgentMessagePostResponseMessage, error) {
agentMessage := decodedAgentMessagePostResponseMessage{
Other: collectOtherKeys(incoming, "responses"),
}
rawResponses, ok := incoming["responses"]
if !ok {
return agentMessage, nil
}
switch responses := rawResponses.(type) {
case nil:
return agentMessage, nil
case []interface{}:
agentMessage.Responses = make([]decodedAgentMessagePostResponse, 0, len(responses))
for i, rawResponse := range responses {
agentMessage.Responses = append(agentMessage.Responses, decodeAgentMessagePostResponse(i, rawResponse))
}
case []map[string]interface{}:
agentMessage.Responses = make([]decodedAgentMessagePostResponse, 0, len(responses))
for i, rawResponse := range responses {
agentMessage.Responses = append(agentMessage.Responses, decodeAgentMessagePostResponse(i, rawResponse))
}
case []agentMessagePostResponse:
agentMessage.Responses = make([]decodedAgentMessagePostResponse, 0, len(responses))
for i, response := range responses {
agentMessage.Responses = append(agentMessage.Responses, decodedAgentMessagePostResponse{
Index: i,
Response: response,
})
}
default:
return agentMessage, fmt.Errorf("responses must be an array, got %T", rawResponses)
}
return agentMessage, nil
}
// decodeAgentMessagePostResponse keeps any agent-owned response tracking keys
// even when the Mythic-owned fields fail to decode.
func decodeAgentMessagePostResponse(index int, rawResponse interface{}) decodedAgentMessagePostResponse {
decodedResponse := decodedAgentMessagePostResponse{
Index: index,
TaskID: extractAgentMessagePostResponseTaskID(rawResponse),
Other: collectAgentMessagePostResponseOtherKeys(rawResponse),
}
if err := mapstructure.Decode(rawResponse, &decodedResponse.Response); err != nil {
decodedResponse.DecodeError = err
}
return decodedResponse
}
func extractAgentMessagePostResponseTaskID(rawResponse interface{}) string {
rawResponseMap, ok := rawResponse.(map[string]interface{})
if !ok {
return ""
}
taskID, ok := rawResponseMap["task_id"].(string)
if !ok {
return ""
}
return taskID
}
func collectAgentMessagePostResponseOtherKeys(rawResponse interface{}) map[string]interface{} {
rawResponseMap, ok := rawResponse.(map[string]interface{})
if !ok {
return map[string]interface{}{}
}
other := make(map[string]interface{}, len(rawResponseMap))
for key, val := range rawResponseMap {
if _, ok := agentMessagePostResponseConsumedKeys[key]; ok {
continue
}
other[key] = val
}
return other
}
func collectValidAgentMessagePostResponses(responses []decodedAgentMessagePostResponse) []agentMessagePostResponse {
validResponses := make([]agentMessagePostResponse, 0, len(responses))
for _, response := range responses {
if response.DecodeError == nil {
validResponses = append(validResponses, response.Response)
}
}
return validResponses
}
func buildAgentMessagePostResponseDecodeError(response decodedAgentMessagePostResponse) map[string]interface{} {
mythicResponse := map[string]interface{}{
"status": "error",
"error": fmt.Sprintf("Failed to decode response[%d]: %s", response.Index, response.DecodeError.Error()),
}
if response.TaskID != "" {
mythicResponse["task_id"] = response.TaskID
}
reflectBackOtherKeys(&mythicResponse, &response.Other)
return mythicResponse
}
func reportAgentMessagePostResponseDecodeErrors(operationID int, responses []decodedAgentMessagePostResponse) {
errorMessages := make([]string, 0)
errorCount := 0
for _, response := range responses {
if response.DecodeError == nil {
continue
}
errorCount++
logging.LogError(response.DecodeError, "Failed to decode individual agent response", "response_index", response.Index)
if len(errorMessages) < 10 {
errorMessages = append(errorMessages, fmt.Sprintf("response[%d]: %s", response.Index, response.DecodeError.Error()))
}
}
if errorCount == 0 {
return
}
if errorCount > len(errorMessages) {
errorMessages = append(errorMessages, "additional response decode errors omitted")
}
go SendAllOperationsMessage(fmt.Sprintf("Failed to decode %d response entr%s from an agent message; valid sibling responses were still processed:\n%s",
errorCount, func() string {
if errorCount == 1 {
return "y"
}
return "ies"
}(), strings.Join(errorMessages, "\n")),
operationID, "agent_message_bad_post_response", database.MESSAGE_LEVEL_AGENT_MESSGAGE, true)
}
func updateAgentMessagePostResponseTask(task databaseStructs.Task) (bool, error) {
tx, err := database.DB.Beginx()
if err != nil {
@@ -499,8 +665,7 @@ func handleAgentMessagePostResponse(incoming *map[string]interface{}, uUIDInfo *
]
}
*/
agentMessage := agentMessagePostResponseMessage{}
err := mapstructure.Decode(incoming, &agentMessage)
agentMessage, err := decodeAgentMessagePostResponseMessage(*incoming)
cachedTaskData := make(map[string]databaseStructs.Task)
cachedFileData := make(map[string]databaseStructs.Filemeta)
if err != nil {
@@ -519,25 +684,32 @@ func handleAgentMessagePostResponse(incoming *map[string]interface{}, uUIDInfo *
return map[string]interface{}{}, err
}
responses := []map[string]interface{}{}
cachedTaskData, err = getAgentMessagePostResponseTasks(agentMessage.Responses)
reportAgentMessagePostResponseDecodeErrors(uUIDInfo.OperationID, agentMessage.Responses)
validAgentResponses := collectValidAgentMessagePostResponses(agentMessage.Responses)
cachedTaskData, err = getAgentMessagePostResponseTasks(validAgentResponses)
if err != nil {
logging.LogError(err, "Failed to batch load tasks for post_response")
}
tasksToUpdate := make(map[string]databaseStructs.Task, len(cachedTaskData))
// iterate over the agent messages
for i, _ := range agentMessage.Responses {
for _, decodedAgentResponse := range agentMessage.Responses {
if decodedAgentResponse.DecodeError != nil {
responses = append(responses, buildAgentMessagePostResponseDecodeError(decodedAgentResponse))
continue
}
agentResponse := decodedAgentResponse.Response
mythicResponse := map[string]interface{}{
"task_id": agentMessage.Responses[i].TaskID,
"task_id": agentResponse.TaskID,
"status": "success",
}
//logging.LogDebug("Got response data from agent", "response data", agentResponse, "extra keys", agentResponse.Other)
// every response should be tied to some task
currentTask, ok := tasksToUpdate[agentMessage.Responses[i].TaskID]
currentTask, ok := tasksToUpdate[agentResponse.TaskID]
if !ok {
currentTask, ok = cachedTaskData[agentMessage.Responses[i].TaskID]
currentTask, ok = cachedTaskData[agentResponse.TaskID]
}
if !ok {
logging.LogError(nil, "Failed to find task", "task id", agentMessage.Responses[i].TaskID)
logging.LogError(nil, "Failed to find task", "task id", agentResponse.TaskID)
mythicResponse["status"] = "error"
mythicResponse["error"] = "Failed to find task"
responses = append(responses, mythicResponse)
@@ -545,26 +717,26 @@ func handleAgentMessagePostResponse(incoming *map[string]interface{}, uUIDInfo *
}
// always process here
if agentMessage.Responses[i].Download != nil {
if agentResponse.Download != nil {
fileMeta := databaseStructs.Filemeta{}
if agentMessage.Responses[i].Download.FileID != nil && *agentMessage.Responses[i].Download.FileID != "" {
if _, ok := cachedFileData[*agentMessage.Responses[i].Download.FileID]; ok {
fileMeta = cachedFileData[*agentMessage.Responses[i].Download.FileID]
if agentResponse.Download.FileID != nil && *agentResponse.Download.FileID != "" {
if _, ok := cachedFileData[*agentResponse.Download.FileID]; ok {
fileMeta = cachedFileData[*agentResponse.Download.FileID]
} else {
fileMeta = databaseStructs.Filemeta{AgentFileID: *agentMessage.Responses[i].Download.FileID}
fileMeta = databaseStructs.Filemeta{AgentFileID: *agentResponse.Download.FileID}
err = database.DB.Get(&fileMeta, `SELECT
id, "path", total_chunks, chunks_received, host, is_screenshot, full_remote_path, complete, md5, sha1, filename, chunk_size, operation_id, mythictree_id, received_chunk_ids
FROM filemeta
WHERE agent_file_id=$1`, *agentMessage.Responses[i].Download.FileID)
WHERE agent_file_id=$1`, *agentResponse.Download.FileID)
if err != nil {
logging.LogError(err, "Failed to find fileID in agent download request", "fileid", *agentMessage.Responses[i].Download.FileID)
logging.LogError(err, "Failed to find fileID in agent download request", "fileid", *agentResponse.Download.FileID)
continue
}
fileMeta.Task = &databaseStructs.Task{}
fileMeta.Task.OperatorID = currentTask.OperatorID
}
}
newFileID, err := handleAgentMessagePostResponseDownload(&currentTask, &agentMessage.Responses[i], &fileMeta)
newFileID, err := handleAgentMessagePostResponseDownload(&currentTask, &agentResponse, &fileMeta)
if err != nil {
mythicResponse["status"] = "error"
mythicResponse["error"] = err.Error()
@@ -572,14 +744,14 @@ func handleAgentMessagePostResponse(incoming *map[string]interface{}, uUIDInfo *
mythicResponse["file_id"] = newFileID
cachedFileData[newFileID] = fileMeta
}
if agentMessage.Responses[i].Download.ChunkNum != nil {
mythicResponse["chunk_num"] = *agentMessage.Responses[i].Download.ChunkNum
if agentResponse.Download.ChunkNum != nil {
mythicResponse["chunk_num"] = *agentResponse.Download.ChunkNum
}
}
// always process here
if agentMessage.Responses[i].Upload != nil {
if uploadResponse, err := handleAgentMessagePostResponseUpload(currentTask, agentMessage.Responses[i]); err != nil {
if agentResponse.Upload != nil {
if uploadResponse, err := handleAgentMessagePostResponseUpload(currentTask, agentResponse); err != nil {
mythicResponse["status"] = "error"
mythicResponse["error"] = err.Error()
logging.LogError(err, "Failed to handle agent upload")
@@ -597,89 +769,89 @@ func handleAgentMessagePostResponse(incoming *map[string]interface{}, uUIDInfo *
currentTask.StatusTimestampProcessed.Valid = true
}
// this section can happen async, but in order
if agentMessage.Responses[i].Completed != nil {
if *agentMessage.Responses[i].Completed {
currentTask.Completed = *agentMessage.Responses[i].Completed
if agentResponse.Completed != nil {
if *agentResponse.Completed {
currentTask.Completed = *agentResponse.Completed
}
}
if agentMessage.Responses[i].Status != nil && *agentMessage.Responses[i].Status != "" {
if agentResponse.Status != nil && *agentResponse.Status != "" {
if currentTask.Status != PT_TASK_FUNCTION_STATUS_COMPLETED {
currentTask.Status = *agentMessage.Responses[i].Status
currentTask.Status = *agentResponse.Status
}
} else if agentMessage.Responses[i].Completed != nil && *agentMessage.Responses[i].Completed {
} else if agentResponse.Completed != nil && *agentResponse.Completed {
currentTask.Status = PT_TASK_FUNCTION_STATUS_COMPLETED
} else if currentTask.Status == PT_TASK_FUNCTION_STATUS_PROCESSING {
currentTask.Status = PT_TASK_FUNCTION_STATUS_PROCESSED
}
if agentMessage.Responses[i].UserOutput != nil && *agentMessage.Responses[i].UserOutput != "" {
if agentResponse.UserOutput != nil && *agentResponse.UserOutput != "" {
// do it in the background - the agent doesn't need the result of this directly
//handleAgentMessagePostResponseUserOutput(currentTask, agentResponse, true)
enqueueAsyncAgentMessagePostResponse(agentMessagePostResponseUserOutputChannelMessage{
Task: currentTask,
Response: *agentMessage.Responses[i].UserOutput,
SequenceNum: agentMessage.Responses[i].SequenceNumber,
Response: *agentResponse.UserOutput,
SequenceNum: agentResponse.SequenceNumber,
})
}
if agentMessage.Responses[i].Stdout != nil {
currentTask.Stdout += *agentMessage.Responses[i].Stdout
if agentResponse.Stdout != nil {
currentTask.Stdout += *agentResponse.Stdout
}
if agentMessage.Responses[i].Stderr != nil {
currentTask.Stderr += *agentMessage.Responses[i].Stderr
if agentResponse.Stderr != nil {
currentTask.Stderr += *agentResponse.Stderr
}
if agentMessage.Responses[i].FileBrowser != nil {
enqueueMythicTreeFileBrowserResponse(currentTask, agentMessage.Responses[i].FileBrowser, 0)
if agentResponse.FileBrowser != nil {
enqueueMythicTreeFileBrowserResponse(currentTask, agentResponse.FileBrowser, 0)
}
if agentMessage.Responses[i].Processes != nil {
enqueueMythicTreeProcessResponse(currentTask, agentMessage.Responses[i].Processes, 0)
if agentResponse.Processes != nil {
enqueueMythicTreeProcessResponse(currentTask, agentResponse.Processes, 0)
}
if agentMessage.Responses[i].RemovedFiles != nil {
go handleAgentMessagePostResponseRemovedFiles(currentTask, agentMessage.Responses[i].RemovedFiles)
if agentResponse.RemovedFiles != nil {
go handleAgentMessagePostResponseRemovedFiles(currentTask, agentResponse.RemovedFiles)
}
if agentMessage.Responses[i].Credentials != nil {
go handleAgentMessagePostResponseCredentials(currentTask, agentMessage.Responses[i].Credentials)
if agentResponse.Credentials != nil {
go handleAgentMessagePostResponseCredentials(currentTask, agentResponse.Credentials)
}
if agentMessage.Responses[i].Keylogs != nil {
go handleAgentMessagePostResponseKeylogs(currentTask, agentMessage.Responses[i].Keylogs)
if agentResponse.Keylogs != nil {
go handleAgentMessagePostResponseKeylogs(currentTask, agentResponse.Keylogs)
}
if agentMessage.Responses[i].Tokens != nil && agentMessage.Responses[i].CallbackTokens != nil {
if agentResponse.Tokens != nil && agentResponse.CallbackTokens != nil {
// need to make sure we process tokens _then_ process callback tokens
go handleAgentMessagePostResponseCallbackTokensAndTokens(currentTask, agentMessage.Responses[i].Tokens, agentMessage.Responses[i].CallbackTokens)
go handleAgentMessagePostResponseCallbackTokensAndTokens(currentTask, agentResponse.Tokens, agentResponse.CallbackTokens)
} else {
if agentMessage.Responses[i].Tokens != nil {
go handleAgentMessagePostResponseTokens(currentTask, agentMessage.Responses[i].Tokens)
if agentResponse.Tokens != nil {
go handleAgentMessagePostResponseTokens(currentTask, agentResponse.Tokens)
}
if agentMessage.Responses[i].CallbackTokens != nil {
go handleAgentMessagePostResponseCallbackTokens(currentTask, agentMessage.Responses[i].CallbackTokens)
if agentResponse.CallbackTokens != nil {
go handleAgentMessagePostResponseCallbackTokens(currentTask, agentResponse.CallbackTokens)
}
}
if agentMessage.Responses[i].ProcessResponse != nil {
go handleAgentMessagePostResponseProcessResponse(currentTask, agentMessage.Responses[i].ProcessResponse)
if agentResponse.ProcessResponse != nil {
go handleAgentMessagePostResponseProcessResponse(currentTask, agentResponse.ProcessResponse)
}
if agentMessage.Responses[i].Commands != nil {
go handleAgentMessagePostResponseCommands(currentTask, agentMessage.Responses[i].Commands)
if agentResponse.Commands != nil {
go handleAgentMessagePostResponseCommands(currentTask, agentResponse.Commands)
}
if agentMessage.Responses[i].Edges != nil {
go handleAgentMessagePostResponseEdges(uUIDInfo, agentMessage.Responses[i].Edges)
if agentResponse.Edges != nil {
go handleAgentMessagePostResponseEdges(uUIDInfo, agentResponse.Edges)
}
if agentMessage.Responses[i].Alerts != nil {
go handleAgentMessagePostResponseAlerts(currentTask.OperationID, uUIDInfo.CallbackID, uUIDInfo.CallbackDisplayID, agentMessage.Responses[i].Alerts)
if agentResponse.Alerts != nil {
go handleAgentMessagePostResponseAlerts(currentTask.OperationID, uUIDInfo.CallbackID, uUIDInfo.CallbackDisplayID, agentResponse.Alerts)
}
if agentMessage.Responses[i].Artifacts != nil {
if agentResponse.Artifacts != nil {
// report back artifact information so that the agent can update the specific artifacts if needed
artifactResponses := handleAgentMessagePostResponseArtifacts(currentTask, agentMessage.Responses[i].Artifacts)
artifactResponses := handleAgentMessagePostResponseArtifacts(currentTask, agentResponse.Artifacts)
mythicResponse["artifacts"] = artifactResponses
}
if agentMessage.Responses[i].Callback != nil {
go handleAgentMessagePostResponseCallback(currentTask, agentMessage.Responses[i].Callback)
if agentResponse.Callback != nil {
go handleAgentMessagePostResponseCallback(currentTask, agentResponse.Callback)
}
if agentMessage.Responses[i].Events != nil && len(*agentMessage.Responses[i].Events) > 0 {
go handleAgentMessagePostResponseEvent(currentTask, agentMessage.Responses[i].Events)
if agentResponse.Events != nil && len(*agentResponse.Events) > 0 {
go handleAgentMessagePostResponseEvent(currentTask, agentResponse.Events)
}
if agentMessage.Responses[i].CustomBrowser != nil {
enqueueMythicTreeCustomBrowserResponse(currentTask, agentMessage.Responses[i].CustomBrowser)
if agentResponse.CustomBrowser != nil {
enqueueMythicTreeCustomBrowserResponse(currentTask, agentResponse.CustomBrowser)
}
// this section always happens
reflectBackOtherKeys(&mythicResponse, &agentMessage.Responses[i].Other)
reflectBackOtherKeys(&mythicResponse, &agentResponse.Other)
responses = append(responses, mythicResponse)
tasksToUpdate[currentTask.AgentTaskID] = currentTask
}
@@ -0,0 +1,172 @@
package rabbitmq
import (
"fmt"
"testing"
"github.com/mitchellh/mapstructure"
)
type benchmarkAgentMessagePostResponseMessage struct {
Responses []agentMessagePostResponse `json:"responses" mapstructure:"responses" xml:"responses"`
Other map[string]interface{} `json:"-" mapstructure:",remain"`
}
var benchmarkDecodedPostResponseMessage decodedAgentMessagePostResponseMessage
var benchmarkWholeArrayPostResponseMessage benchmarkAgentMessagePostResponseMessage
func TestDecodeAgentMessagePostResponseMessageIsolatesMalformedResponses(t *testing.T) {
incoming := map[string]interface{}{
"action": "post_response",
"batch_tracking": "batch-1",
"responses": []interface{}{
map[string]interface{}{
"task_id": "task-good",
"user_output": "good output",
"response_tracking": "response-1",
},
map[string]interface{}{
"task_id": "task-bad",
"download": "not a download object",
"response_tracking": "response-2",
},
},
}
decodedMessage, err := decodeAgentMessagePostResponseMessage(incoming)
if err != nil {
t.Fatalf("unexpected top-level decode error: %v", err)
}
if decodedMessage.Other["batch_tracking"] != "batch-1" {
t.Fatalf("expected top-level tracking key to be preserved, got %#v", decodedMessage.Other)
}
if len(decodedMessage.Responses) != 2 {
t.Fatalf("expected two decoded response entries, got %d", len(decodedMessage.Responses))
}
if decodedMessage.Responses[0].DecodeError != nil {
t.Fatalf("expected first response to decode successfully, got %v", decodedMessage.Responses[0].DecodeError)
}
if decodedMessage.Responses[0].Response.Other["response_tracking"] != "response-1" {
t.Fatalf("expected response-level tracking key to be preserved, got %#v", decodedMessage.Responses[0].Response.Other)
}
if decodedMessage.Responses[1].DecodeError == nil {
t.Fatal("expected malformed response to have a decode error")
}
if decodedMessage.Responses[1].TaskID != "task-bad" {
t.Fatalf("expected task_id to be recovered from malformed response, got %q", decodedMessage.Responses[1].TaskID)
}
if decodedMessage.Responses[1].Other["response_tracking"] != "response-2" {
t.Fatalf("expected malformed response tracking key to be preserved, got %#v", decodedMessage.Responses[1].Other)
}
if _, ok := decodedMessage.Responses[1].Other["download"]; ok {
t.Fatalf("expected consumed response key to be omitted from reflected keys, got %#v", decodedMessage.Responses[1].Other)
}
mythicResponse := map[string]interface{}{
"task_id": decodedMessage.Responses[0].Response.TaskID,
"status": "success",
}
reflectBackOtherKeys(&mythicResponse, &decodedMessage.Responses[0].Response.Other)
if mythicResponse["response_tracking"] != "response-1" {
t.Fatalf("expected successful response tracking key to be reflected, got %#v", mythicResponse)
}
badMythicResponse := buildAgentMessagePostResponseDecodeError(decodedMessage.Responses[1])
if badMythicResponse["task_id"] != "task-bad" {
t.Fatalf("expected malformed response task_id to be reflected in error response, got %#v", badMythicResponse)
}
if badMythicResponse["response_tracking"] != "response-2" {
t.Fatalf("expected malformed response tracking key to be reflected, got %#v", badMythicResponse)
}
if _, ok := badMythicResponse["download"]; ok {
t.Fatalf("expected consumed malformed response key to stay out of error response, got %#v", badMythicResponse)
}
}
func TestHandleAgentMessageGetTaskingReflectsOnlyUnusedKeys(t *testing.T) {
replaceSubmittedTasksForProxyTest(t, nil)
incoming := map[string]interface{}{
"action": "get_tasking",
"tasking_size": float64(-1),
"get_delegate_tasks": false,
"agent_tracking": "track-me",
}
response, err := handleAgentMessageGetTasking(&incoming, 1)
if err != nil {
t.Fatalf("unexpected get_tasking error: %v", err)
}
if response["agent_tracking"] != "track-me" {
t.Fatalf("expected custom get_tasking tracking key to be reflected, got %#v", response)
}
if _, ok := response["tasking_size"]; ok {
t.Fatalf("expected consumed tasking_size key to be omitted from response, got %#v", response)
}
if _, ok := response["get_delegate_tasks"]; ok {
t.Fatalf("expected consumed get_delegate_tasks key to be omitted from response, got %#v", response)
}
if _, ok := incoming["tasking_size"]; ok {
t.Fatalf("expected tasking_size to be removed after get_tasking processing, got %#v", incoming)
}
if incoming["get_delegate_tasks"] != false {
t.Fatalf("expected get_delegate_tasks to remain for outer delegate processing, got %#v", incoming)
}
}
func BenchmarkDecodeAgentMessagePostResponseMessage(b *testing.B) {
benchmarks := []struct {
name string
responseCount int
badIndex int
}{
{name: "valid_1", responseCount: 1, badIndex: -1},
{name: "valid_10", responseCount: 10, badIndex: -1},
{name: "valid_100", responseCount: 100, badIndex: -1},
{name: "one_bad_10", responseCount: 10, badIndex: 5},
{name: "one_bad_100", responseCount: 100, badIndex: 50},
}
for _, benchmark := range benchmarks {
incoming := buildBenchmarkPostResponseMessage(benchmark.responseCount, benchmark.badIndex)
b.Run("whole_array_mapstructure/"+benchmark.name, func(b *testing.B) {
b.ReportAllocs()
for i := 0; i < b.N; i++ {
var decoded benchmarkAgentMessagePostResponseMessage
_ = mapstructure.Decode(incoming, &decoded)
benchmarkWholeArrayPostResponseMessage = decoded
}
})
b.Run("per_response_decode/"+benchmark.name, func(b *testing.B) {
b.ReportAllocs()
for i := 0; i < b.N; i++ {
decoded, _ := decodeAgentMessagePostResponseMessage(incoming)
benchmarkDecodedPostResponseMessage = decoded
}
})
}
}
func buildBenchmarkPostResponseMessage(responseCount int, badIndex int) map[string]interface{} {
responses := make([]interface{}, 0, responseCount)
for i := 0; i < responseCount; i++ {
response := map[string]interface{}{
"task_id": fmt.Sprintf("task-%d", i),
"completed": i%3 == 0,
"status": "processed",
"user_output": fmt.Sprintf("output-%d", i),
"stdout": fmt.Sprintf("stdout-%d", i),
"stderr": fmt.Sprintf("stderr-%d", i),
"sequence_num": int64(i),
"response_tracking": fmt.Sprintf("response-tracking-%d", i),
}
if i == badIndex {
response["download"] = "not a download object"
}
responses = append(responses, response)
}
return map[string]interface{}{
"action": "post_response",
"batch_tracking": "batch-tracking",
"responses": responses,
}
}