mirror of
https://github.com/its-a-feature/Mythic
synced 2026-06-08 14:55:38 +00:00
per response decode to not bulk bail on bad entries
This commit is contained in:
@@ -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(¤tTask, &agentMessage.Responses[i], &fileMeta)
|
||||
newFileID, err := handleAgentMessagePostResponseDownload(¤tTask, &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,
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user