#include "source/fuzz/transformation_composite_construct.h"
#include "source/fuzz/data_descriptor.h"
#include "source/fuzz/fuzzer_util.h"
#include "source/fuzz/instruction_descriptor.h"
#include "source/opt/instruction.h"
namespace spvtools {
namespace fuzz {
TransformationCompositeConstruct::TransformationCompositeConstruct(
protobufs::TransformationCompositeConstruct message)
: message_(std::move(message)) {}
TransformationCompositeConstruct::TransformationCompositeConstruct(
uint32_t composite_type_id, std::vector<uint32_t> component,
const protobufs::InstructionDescriptor& instruction_to_insert_before,
uint32_t fresh_id) {
message_.set_composite_type_id(composite_type_id);
for (auto a_component : component) {
message_.add_component(a_component);
}
*message_.mutable_instruction_to_insert_before() =
instruction_to_insert_before;
message_.set_fresh_id(fresh_id);
}
bool TransformationCompositeConstruct::IsApplicable(
opt::IRContext* ir_context, const TransformationContext& ) const {
if (!fuzzerutil::IsFreshId(ir_context, message_.fresh_id())) {
return false;
}
auto insert_before =
FindInstruction(message_.instruction_to_insert_before(), ir_context);
if (!insert_before) {
return false;
}
auto composite_type =
ir_context->get_type_mgr()->GetType(message_.composite_type_id());
if (!fuzzerutil::IsCompositeType(composite_type)) {
return false;
}
if (composite_type->AsArray() &&
!ComponentsForArrayConstructionAreOK(ir_context,
*composite_type->AsArray())) {
return false;
}
if (composite_type->AsMatrix() &&
!ComponentsForMatrixConstructionAreOK(ir_context,
*composite_type->AsMatrix())) {
return false;
}
if (composite_type->AsStruct() &&
!ComponentsForStructConstructionAreOK(ir_context,
*composite_type->AsStruct())) {
return false;
}
if (composite_type->AsVector() &&
!ComponentsForVectorConstructionAreOK(ir_context,
*composite_type->AsVector())) {
return false;
}
for (auto component : message_.component()) {
auto* inst = ir_context->get_def_use_mgr()->GetDef(component);
if (!inst) {
return false;
}
if (!fuzzerutil::IdIsAvailableBeforeInstruction(ir_context, insert_before,
component)) {
return false;
}
}
return true;
}
void TransformationCompositeConstruct::Apply(
opt::IRContext* ir_context,
TransformationContext* transformation_context) const {
auto insert_before_inst =
FindInstruction(message_.instruction_to_insert_before(), ir_context);
auto destination_block = ir_context->get_instr_block(insert_before_inst);
auto insert_before = fuzzerutil::GetIteratorForInstruction(
destination_block, insert_before_inst);
opt::Instruction::OperandList in_operands;
for (auto& component_id : message_.component()) {
in_operands.push_back({SPV_OPERAND_TYPE_ID, {component_id}});
}
auto new_instruction = MakeUnique<opt::Instruction>(
ir_context, spv::Op::OpCompositeConstruct, message_.composite_type_id(),
message_.fresh_id(), in_operands);
auto new_instruction_ptr = new_instruction.get();
insert_before.InsertBefore(std::move(new_instruction));
ir_context->get_def_use_mgr()->AnalyzeInstDefUse(new_instruction_ptr);
ir_context->set_instr_block(new_instruction_ptr, destination_block);
fuzzerutil::UpdateModuleIdBound(ir_context, message_.fresh_id());
AddDataSynonymFacts(ir_context, transformation_context);
}
bool TransformationCompositeConstruct::ComponentsForArrayConstructionAreOK(
opt::IRContext* ir_context, const opt::analysis::Array& array_type) const {
if (array_type.length_info().words[0] !=
opt::analysis::Array::LengthInfo::kConstant) {
return false;
}
if (array_type.length_info().words.size() != 2) {
return false;
}
auto array_size = array_type.length_info().words[1];
if (static_cast<uint32_t>(message_.component().size()) != array_size) {
return false;
}
for (auto component_id : message_.component()) {
auto inst = ir_context->get_def_use_mgr()->GetDef(component_id);
if (inst == nullptr || !inst->type_id()) {
return false;
}
auto component_type = ir_context->get_type_mgr()->GetType(inst->type_id());
assert(component_type);
if (component_type != array_type.element_type()) {
return false;
}
}
return true;
}
bool TransformationCompositeConstruct::ComponentsForMatrixConstructionAreOK(
opt::IRContext* ir_context,
const opt::analysis::Matrix& matrix_type) const {
if (static_cast<uint32_t>(message_.component().size()) !=
matrix_type.element_count()) {
return false;
}
for (auto component_id : message_.component()) {
auto inst = ir_context->get_def_use_mgr()->GetDef(component_id);
if (inst == nullptr || !inst->type_id()) {
return false;
}
auto component_type = ir_context->get_type_mgr()->GetType(inst->type_id());
assert(component_type);
if (component_type != matrix_type.element_type()) {
return false;
}
}
return true;
}
bool TransformationCompositeConstruct::ComponentsForStructConstructionAreOK(
opt::IRContext* ir_context,
const opt::analysis::Struct& struct_type) const {
if (static_cast<uint32_t>(message_.component().size()) !=
struct_type.element_types().size()) {
return false;
}
for (uint32_t field_index = 0;
field_index < struct_type.element_types().size(); field_index++) {
auto inst = ir_context->get_def_use_mgr()->GetDef(
message_.component()[field_index]);
if (inst == nullptr || !inst->type_id()) {
return false;
}
auto component_type = ir_context->get_type_mgr()->GetType(inst->type_id());
assert(component_type);
if (component_type != struct_type.element_types()[field_index]) {
return false;
}
}
return true;
}
bool TransformationCompositeConstruct::ComponentsForVectorConstructionAreOK(
opt::IRContext* ir_context,
const opt::analysis::Vector& vector_type) const {
uint32_t base_element_count = 0;
auto element_type = vector_type.element_type();
for (auto& component_id : message_.component()) {
auto inst = ir_context->get_def_use_mgr()->GetDef(component_id);
if (inst == nullptr || !inst->type_id()) {
return false;
}
auto component_type = ir_context->get_type_mgr()->GetType(inst->type_id());
assert(component_type);
if (component_type == element_type) {
base_element_count++;
} else if (component_type->AsVector() &&
component_type->AsVector()->element_type() == element_type) {
base_element_count += component_type->AsVector()->element_count();
} else {
return false;
}
}
return base_element_count == vector_type.element_count();
}
protobufs::Transformation TransformationCompositeConstruct::ToMessage() const {
protobufs::Transformation result;
*result.mutable_composite_construct() = message_;
return result;
}
std::unordered_set<uint32_t> TransformationCompositeConstruct::GetFreshIds()
const {
return {message_.fresh_id()};
}
void TransformationCompositeConstruct::AddDataSynonymFacts(
opt::IRContext* ir_context,
TransformationContext* transformation_context) const {
if (transformation_context->GetFactManager()->IdIsIrrelevant(
message_.fresh_id())) {
return;
}
auto composite_type =
ir_context->get_type_mgr()->GetType(message_.composite_type_id());
uint32_t index = 0;
for (auto component : message_.component()) {
auto component_type = ir_context->get_type_mgr()->GetType(
ir_context->get_def_use_mgr()->GetDef(component)->type_id());
const bool packing_vector_into_vector =
composite_type->AsVector() && component_type->AsVector();
if (!fuzzerutil::CanMakeSynonymOf(
ir_context, *transformation_context,
*ir_context->get_def_use_mgr()->GetDef(component))) {
index += packing_vector_into_vector
? component_type->AsVector()->element_count()
: 1;
continue;
}
if (packing_vector_into_vector) {
assert(component_type->AsVector()->element_type() ==
composite_type->AsVector()->element_type());
assert(component_type->AsVector()->element_count() <
composite_type->AsVector()->element_count());
for (uint32_t subvector_index = 0;
subvector_index < component_type->AsVector()->element_count();
subvector_index++) {
transformation_context->GetFactManager()->AddFactDataSynonym(
MakeDataDescriptor(component, {subvector_index}),
MakeDataDescriptor(message_.fresh_id(), {index}));
index++;
}
} else {
transformation_context->GetFactManager()->AddFactDataSynonym(
MakeDataDescriptor(component, {}),
MakeDataDescriptor(message_.fresh_id(), {index}));
index++;
}
}
}
}
}