#include "source/fuzz/transformation_vector_shuffle.h"
#include "source/fuzz/fuzzer_util.h"
#include "source/fuzz/instruction_descriptor.h"
namespace spvtools {
namespace fuzz {
TransformationVectorShuffle::TransformationVectorShuffle(
protobufs::TransformationVectorShuffle message)
: message_(std::move(message)) {}
TransformationVectorShuffle::TransformationVectorShuffle(
const protobufs::InstructionDescriptor& instruction_to_insert_before,
uint32_t fresh_id, uint32_t vector1, uint32_t vector2,
const std::vector<uint32_t>& component) {
*message_.mutable_instruction_to_insert_before() =
instruction_to_insert_before;
message_.set_fresh_id(fresh_id);
message_.set_vector1(vector1);
message_.set_vector2(vector2);
for (auto a_component : component) {
message_.add_component(a_component);
}
}
bool TransformationVectorShuffle::IsApplicable(
opt::IRContext* ir_context, const TransformationContext& ) const {
if (!fuzzerutil::IsFreshId(ir_context, message_.fresh_id())) {
return false;
}
auto instruction_to_insert_before =
FindInstruction(message_.instruction_to_insert_before(), ir_context);
if (!instruction_to_insert_before) {
return false;
}
auto vector1_instruction =
ir_context->get_def_use_mgr()->GetDef(message_.vector1());
if (!vector1_instruction || !vector1_instruction->type_id()) {
return false;
}
auto vector2_instruction =
ir_context->get_def_use_mgr()->GetDef(message_.vector2());
if (!vector2_instruction || !vector2_instruction->type_id()) {
return false;
}
auto vector1_type =
ir_context->get_type_mgr()->GetType(vector1_instruction->type_id());
if (!vector1_type->AsVector()) {
return false;
}
auto vector2_type =
ir_context->get_type_mgr()->GetType(vector2_instruction->type_id());
if (!vector2_type->AsVector()) {
return false;
}
if (vector1_type->AsVector()->element_type() !=
vector2_type->AsVector()->element_type()) {
return false;
}
uint32_t combined_size = vector1_type->AsVector()->element_count() +
vector2_type->AsVector()->element_count();
for (auto a_compoment : message_.component()) {
if (a_compoment != 0xFFFFFFFF && a_compoment >= combined_size) {
return false;
}
}
if (!GetResultTypeId(ir_context, *vector1_type->AsVector()->element_type())) {
return false;
}
for (auto used_instruction : {vector1_instruction, vector2_instruction}) {
if (auto block = ir_context->get_instr_block(used_instruction)) {
if (!ir_context->GetDominatorAnalysis(block->GetParent())
->Dominates(used_instruction, instruction_to_insert_before)) {
return false;
}
}
}
return fuzzerutil::CanInsertOpcodeBeforeInstruction(
spv::Op::OpVectorShuffle, instruction_to_insert_before);
}
void TransformationVectorShuffle::Apply(
opt::IRContext* ir_context,
TransformationContext* transformation_context) const {
opt::Instruction::OperandList shuffle_operands = {
{SPV_OPERAND_TYPE_ID, {message_.vector1()}},
{SPV_OPERAND_TYPE_ID, {message_.vector2()}}};
for (auto a_component : message_.component()) {
shuffle_operands.push_back(
{SPV_OPERAND_TYPE_LITERAL_INTEGER, {a_component}});
}
uint32_t result_type_id = GetResultTypeId(
ir_context,
*GetVectorType(ir_context, message_.vector1())->element_type());
auto insert_before =
FindInstruction(message_.instruction_to_insert_before(), ir_context);
opt::Instruction* new_instruction =
insert_before->InsertBefore(MakeUnique<opt::Instruction>(
ir_context, spv::Op::OpVectorShuffle, result_type_id,
message_.fresh_id(), shuffle_operands));
fuzzerutil::UpdateModuleIdBound(ir_context, message_.fresh_id());
ir_context->get_def_use_mgr()->AnalyzeInstDefUse(new_instruction);
ir_context->set_instr_block(new_instruction,
ir_context->get_instr_block(insert_before));
AddDataSynonymFacts(ir_context, transformation_context);
}
protobufs::Transformation TransformationVectorShuffle::ToMessage() const {
protobufs::Transformation result;
*result.mutable_vector_shuffle() = message_;
return result;
}
uint32_t TransformationVectorShuffle::GetResultTypeId(
opt::IRContext* ir_context, const opt::analysis::Type& element_type) const {
opt::analysis::Vector result_type(
&element_type, static_cast<uint32_t>(message_.component_size()));
return ir_context->get_type_mgr()->GetId(&result_type);
}
opt::analysis::Vector* TransformationVectorShuffle::GetVectorType(
opt::IRContext* ir_context, uint32_t id_of_vector) {
return ir_context->get_type_mgr()
->GetType(ir_context->get_def_use_mgr()->GetDef(id_of_vector)->type_id())
->AsVector();
}
std::unordered_set<uint32_t> TransformationVectorShuffle::GetFreshIds() const {
return {message_.fresh_id()};
}
void TransformationVectorShuffle::AddDataSynonymFacts(
opt::IRContext* ir_context,
TransformationContext* transformation_context) const {
if (transformation_context->GetFactManager()->IdIsIrrelevant(
message_.fresh_id())) {
return;
}
for (uint32_t component_index = 0;
component_index < static_cast<uint32_t>(message_.component_size());
component_index++) {
uint32_t component = message_.component(component_index);
if (component == 0xFFFFFFFF) {
continue;
}
protobufs::DataDescriptor descriptor_for_result_component =
MakeDataDescriptor(message_.fresh_id(), {component_index});
protobufs::DataDescriptor descriptor_for_source_component;
if (component <
GetVectorType(ir_context, message_.vector1())->element_count()) {
if (!fuzzerutil::CanMakeSynonymOf(
ir_context, *transformation_context,
*ir_context->get_def_use_mgr()->GetDef(message_.vector1()))) {
continue;
}
descriptor_for_source_component =
MakeDataDescriptor(message_.vector1(), {component});
} else {
if (!fuzzerutil::CanMakeSynonymOf(
ir_context, *transformation_context,
*ir_context->get_def_use_mgr()->GetDef(message_.vector2()))) {
continue;
}
auto index_into_vector_2 =
component -
GetVectorType(ir_context, message_.vector1())->element_count();
assert(
index_into_vector_2 <
GetVectorType(ir_context, message_.vector2())->element_count() &&
"Vector shuffle index is out of bounds.");
descriptor_for_source_component =
MakeDataDescriptor(message_.vector2(), {index_into_vector_2});
}
transformation_context->GetFactManager()->AddFactDataSynonym(
descriptor_for_result_component, descriptor_for_source_component);
}
}
}
}