#include "tools/cfg/bin_to_dot.h"
#include <cassert>
#include <iostream>
#include <utility>
#include <vector>
#include "source/assembly_grammar.h"
#include "source/name_mapper.h"
namespace {
const char* kMergeStyle = "style=dashed";
const char* kContinueStyle = "style=dotted";
class DotConverter {
public:
DotConverter(spvtools::NameMapper name_mapper, std::iostream* out)
: name_mapper_(std::move(name_mapper)), out_(*out) {}
void Begin() const {
out_ << "digraph {\n";
out_ << "legend_merge_src [shape=plaintext, label=\"\"];\n"
<< "legend_merge_dest [shape=plaintext, label=\"\"];\n"
<< "legend_merge_src -> legend_merge_dest [label=\" merge\","
<< kMergeStyle << "];\n"
<< "legend_continue_src [shape=plaintext, label=\"\"];\n"
<< "legend_continue_dest [shape=plaintext, label=\"\"];\n"
<< "legend_continue_src -> legend_continue_dest [label=\" continue\","
<< kContinueStyle << "];\n";
}
void End() const { out_ << "}\n"; }
spv_result_t HandleInstruction(const spv_parsed_instruction_t& inst);
private:
void FlushBlock(const std::vector<uint32_t>& successors);
uint32_t current_function_id_ = 0;
uint32_t current_block_id_ = 0;
bool seen_function_entry_block_ = false;
uint32_t merge_ = 0;
uint32_t continue_target_ = 0;
spvtools::NameMapper name_mapper_;
std::ostream& out_;
};
spv_result_t DotConverter::HandleInstruction(
const spv_parsed_instruction_t& inst) {
switch (spv::Op(inst.opcode)) {
case spv::Op::OpFunction:
current_function_id_ = inst.result_id;
seen_function_entry_block_ = false;
break;
case spv::Op::OpFunctionEnd:
current_function_id_ = 0;
break;
case spv::Op::OpLabel:
current_block_id_ = inst.result_id;
break;
case spv::Op::OpBranch:
FlushBlock({inst.words[1]});
break;
case spv::Op::OpBranchConditional:
FlushBlock({inst.words[2], inst.words[3]});
break;
case spv::Op::OpSwitch: {
std::vector<uint32_t> successors{inst.words[2]};
for (size_t i = 3; i < inst.num_operands; i += 2) {
successors.push_back(inst.words[inst.operands[i].offset]);
}
FlushBlock(successors);
} break;
case spv::Op::OpKill:
case spv::Op::OpReturn:
case spv::Op::OpUnreachable:
case spv::Op::OpReturnValue:
FlushBlock({});
break;
case spv::Op::OpLoopMerge:
merge_ = inst.words[1];
continue_target_ = inst.words[2];
break;
case spv::Op::OpSelectionMerge:
merge_ = inst.words[1];
break;
default:
break;
}
return SPV_SUCCESS;
}
void DotConverter::FlushBlock(const std::vector<uint32_t>& successors) {
out_ << current_block_id_;
if (!seen_function_entry_block_) {
out_ << " [label=\"" << name_mapper_(current_block_id_) << "\nFn "
<< name_mapper_(current_function_id_) << " entry\", shape=box];\n";
} else {
out_ << " [label=\"" << name_mapper_(current_block_id_) << "\"];\n";
}
for (auto successor : successors) {
out_ << current_block_id_ << " -> " << successor << ";\n";
}
if (merge_) {
out_ << current_block_id_ << " -> " << merge_ << " [" << kMergeStyle
<< "];\n";
}
if (continue_target_) {
out_ << current_block_id_ << " -> " << continue_target_ << " ["
<< kContinueStyle << "];\n";
}
seen_function_entry_block_ = true;
merge_ = 0;
continue_target_ = 0;
}
spv_result_t HandleInstruction(
void* user_data, const spv_parsed_instruction_t* parsed_instruction) {
assert(user_data);
auto converter = static_cast<DotConverter*>(user_data);
return converter->HandleInstruction(*parsed_instruction);
}
}
spv_result_t BinaryToDot(const spv_const_context context, const uint32_t* words,
size_t num_words, std::iostream* out,
spv_diagnostic* diagnostic) {
if (!diagnostic) return SPV_ERROR_INVALID_DIAGNOSTIC;
const spvtools::AssemblyGrammar grammar(context);
if (!grammar.isValid()) return SPV_ERROR_INVALID_TABLE;
spvtools::FriendlyNameMapper friendly_mapper(context, words, num_words);
DotConverter converter(friendly_mapper.GetNameMapper(), out);
converter.Begin();
if (auto error = spvBinaryParse(context, &converter, words, num_words,
nullptr, HandleInstruction, diagnostic)) {
return error;
}
converter.End();
return SPV_SUCCESS;
}