#include "source/fuzz/transformation_outline_function.h"
#include <set>
#include "source/fuzz/fuzzer_util.h"
namespace spvtools {
namespace fuzz {
TransformationOutlineFunction::TransformationOutlineFunction(
protobufs::TransformationOutlineFunction message)
: message_(std::move(message)) {}
TransformationOutlineFunction::TransformationOutlineFunction(
uint32_t entry_block, uint32_t exit_block,
uint32_t new_function_struct_return_type_id, uint32_t new_function_type_id,
uint32_t new_function_id, uint32_t new_function_region_entry_block,
uint32_t new_caller_result_id, uint32_t new_callee_result_id,
const std::map<uint32_t, uint32_t>& input_id_to_fresh_id,
const std::map<uint32_t, uint32_t>& output_id_to_fresh_id) {
message_.set_entry_block(entry_block);
message_.set_exit_block(exit_block);
message_.set_new_function_struct_return_type_id(
new_function_struct_return_type_id);
message_.set_new_function_type_id(new_function_type_id);
message_.set_new_function_id(new_function_id);
message_.set_new_function_region_entry_block(new_function_region_entry_block);
message_.set_new_caller_result_id(new_caller_result_id);
message_.set_new_callee_result_id(new_callee_result_id);
*message_.mutable_input_id_to_fresh_id() =
fuzzerutil::MapToRepeatedUInt32Pair(input_id_to_fresh_id);
*message_.mutable_output_id_to_fresh_id() =
fuzzerutil::MapToRepeatedUInt32Pair(output_id_to_fresh_id);
}
bool TransformationOutlineFunction::IsApplicable(
opt::IRContext* ir_context,
const TransformationContext& transformation_context) const {
std::set<uint32_t> ids_used_by_this_transformation;
if (!CheckIdIsFreshAndNotUsedByThisTransformation(
message_.new_function_struct_return_type_id(), ir_context,
&ids_used_by_this_transformation)) {
return false;
}
if (!CheckIdIsFreshAndNotUsedByThisTransformation(
message_.new_function_type_id(), ir_context,
&ids_used_by_this_transformation)) {
return false;
}
if (!CheckIdIsFreshAndNotUsedByThisTransformation(
message_.new_function_id(), ir_context,
&ids_used_by_this_transformation)) {
return false;
}
if (!CheckIdIsFreshAndNotUsedByThisTransformation(
message_.new_function_region_entry_block(), ir_context,
&ids_used_by_this_transformation)) {
return false;
}
if (!CheckIdIsFreshAndNotUsedByThisTransformation(
message_.new_caller_result_id(), ir_context,
&ids_used_by_this_transformation)) {
return false;
}
if (!CheckIdIsFreshAndNotUsedByThisTransformation(
message_.new_callee_result_id(), ir_context,
&ids_used_by_this_transformation)) {
return false;
}
for (auto& pair : message_.input_id_to_fresh_id()) {
if (!CheckIdIsFreshAndNotUsedByThisTransformation(
pair.second(), ir_context, &ids_used_by_this_transformation)) {
return false;
}
}
for (auto& pair : message_.output_id_to_fresh_id()) {
if (!CheckIdIsFreshAndNotUsedByThisTransformation(
pair.second(), ir_context, &ids_used_by_this_transformation)) {
return false;
}
}
for (auto block_id : {message_.entry_block(), message_.exit_block()}) {
auto block_label = ir_context->get_def_use_mgr()->GetDef(block_id);
if (!block_label || block_label->opcode() != spv::Op::OpLabel) {
return false;
}
}
auto entry_block = ir_context->cfg()->block(message_.entry_block());
auto exit_block = ir_context->cfg()->block(message_.exit_block());
if (entry_block->begin()->opcode() == spv::Op::OpVariable) {
return false;
}
if (entry_block->GetLoopMergeInst()) {
return false;
}
if (fuzzerutil::IsMergeOrContinue(ir_context, exit_block->id())) {
return false;
}
if (entry_block->begin()->opcode() == spv::Op::OpPhi) {
return false;
}
if (entry_block->GetParent() != exit_block->GetParent()) {
return false;
}
auto dominator_analysis =
ir_context->GetDominatorAnalysis(entry_block->GetParent());
if (!dominator_analysis->Dominates(entry_block, exit_block)) {
return false;
}
auto postdominator_analysis =
ir_context->GetPostDominatorAnalysis(entry_block->GetParent());
if (!postdominator_analysis->Dominates(exit_block, entry_block)) {
return false;
}
auto region_set = GetRegionBlocks(
ir_context,
entry_block = ir_context->cfg()->block(message_.entry_block()),
exit_block = ir_context->cfg()->block(message_.exit_block()));
for (auto& block : *entry_block->GetParent()) {
if (region_set.count(&block) != 0) {
for (auto pred : ir_context->cfg()->preds(block.id())) {
if (!ir_context->IsReachable(*ir_context->cfg()->block(pred))) {
return false;
}
}
}
if (&block == exit_block) {
if (block.GetLoopMergeInst()) {
return false;
}
continue;
}
if (region_set.count(&block) != 0) {
bool all_successors_in_region = true;
block.WhileEachSuccessorLabel([&all_successors_in_region, ir_context,
®ion_set](uint32_t successor) -> bool {
if (region_set.count(ir_context->cfg()->block(successor)) == 0) {
all_successors_in_region = false;
return false;
}
return true;
});
if (!all_successors_in_region) {
return false;
}
}
if (auto merge = block.GetMergeInst()) {
auto merge_block =
ir_context->cfg()->block(merge->GetSingleWordOperand(0));
if (region_set.count(&block) != region_set.count(merge_block)) {
return false;
}
}
if (auto loop_merge = block.GetLoopMergeInst()) {
auto continue_target =
ir_context->cfg()->block(loop_merge->GetSingleWordOperand(1));
if (continue_target != exit_block &&
region_set.count(&block) != region_set.count(continue_target)) {
return false;
}
}
}
auto input_id_to_fresh_id_map =
fuzzerutil::RepeatedUInt32PairToMap(message_.input_id_to_fresh_id());
for (auto id : GetRegionInputIds(ir_context, region_set, exit_block)) {
if (input_id_to_fresh_id_map.count(id) == 0 &&
!transformation_context.GetOverflowIdSource()->HasOverflowIds()) {
return false;
}
auto input_id_inst = ir_context->get_def_use_mgr()->GetDef(id);
if (ir_context->get_def_use_mgr()
->GetDef(input_id_inst->type_id())
->opcode() == spv::Op::OpTypePointer) {
switch (input_id_inst->opcode()) {
case spv::Op::OpFunctionParameter:
case spv::Op::OpVariable:
break;
default:
return false;
}
}
}
auto output_id_to_fresh_id_map =
fuzzerutil::RepeatedUInt32PairToMap(message_.output_id_to_fresh_id());
for (auto id : GetRegionOutputIds(ir_context, region_set, exit_block)) {
if (
(output_id_to_fresh_id_map.count(id) == 0 &&
!transformation_context.GetOverflowIdSource()->HasOverflowIds())
|| ir_context->get_def_use_mgr()
->GetDef(fuzzerutil::GetTypeId(ir_context, id))
->opcode() == spv::Op::OpTypePointer) {
return false;
}
}
return true;
}
void TransformationOutlineFunction::Apply(
opt::IRContext* ir_context,
TransformationContext* transformation_context) const {
auto original_region_entry_block =
ir_context->cfg()->block(message_.entry_block());
auto original_region_exit_block =
ir_context->cfg()->block(message_.exit_block());
std::set<opt::BasicBlock*> region_blocks = GetRegionBlocks(
ir_context, original_region_entry_block, original_region_exit_block);
std::vector<uint32_t> region_input_ids =
GetRegionInputIds(ir_context, region_blocks, original_region_exit_block);
std::vector<uint32_t> region_output_ids =
GetRegionOutputIds(ir_context, region_blocks, original_region_exit_block);
auto input_id_to_fresh_id_map =
fuzzerutil::RepeatedUInt32PairToMap(message_.input_id_to_fresh_id());
auto output_id_to_fresh_id_map =
fuzzerutil::RepeatedUInt32PairToMap(message_.output_id_to_fresh_id());
for (uint32_t id : region_input_ids) {
if (input_id_to_fresh_id_map.count(id) == 0) {
input_id_to_fresh_id_map.insert(
{id,
transformation_context->GetOverflowIdSource()->GetNextOverflowId()});
}
}
for (uint32_t id : region_output_ids) {
if (output_id_to_fresh_id_map.count(id) == 0) {
output_id_to_fresh_id_map.insert(
{id,
transformation_context->GetOverflowIdSource()->GetNextOverflowId()});
}
}
UpdateModuleIdBoundForFreshIds(ir_context, input_id_to_fresh_id_map,
output_id_to_fresh_id_map);
std::map<uint32_t, uint32_t> output_id_to_type_id;
for (uint32_t output_id : region_output_ids) {
output_id_to_type_id[output_id] =
ir_context->get_def_use_mgr()->GetDef(output_id)->type_id();
}
std::unique_ptr<opt::Instruction> cloned_exit_block_terminator =
std::unique_ptr<opt::Instruction>(
original_region_exit_block->terminator()->Clone(ir_context));
std::unique_ptr<opt::Instruction> cloned_exit_block_merge =
original_region_exit_block->GetMergeInst()
? std::unique_ptr<opt::Instruction>(
original_region_exit_block->GetMergeInst()->Clone(ir_context))
: nullptr;
std::unique_ptr<opt::Function> outlined_function = PrepareFunctionPrototype(
region_input_ids, region_output_ids, input_id_to_fresh_id_map, ir_context,
transformation_context);
RemapInputAndOutputIdsInRegion(
ir_context, *original_region_exit_block, region_blocks, region_input_ids,
region_output_ids, input_id_to_fresh_id_map, output_id_to_fresh_id_map);
PopulateOutlinedFunction(
*original_region_entry_block, *original_region_exit_block, region_blocks,
region_output_ids, output_id_to_type_id, output_id_to_fresh_id_map,
ir_context, outlined_function.get());
ShrinkOriginalRegion(
ir_context, region_blocks, region_input_ids, region_output_ids,
output_id_to_type_id, outlined_function->type_id(),
std::move(cloned_exit_block_merge),
std::move(cloned_exit_block_terminator), original_region_entry_block);
const auto* outlined_function_ptr = outlined_function.get();
ir_context->module()->AddFunction(std::move(outlined_function));
ir_context->InvalidateAnalysesExceptFor(
opt::IRContext::Analysis::kAnalysisNone);
if (transformation_context->GetFactManager()->FunctionIsLivesafe(
original_region_entry_block->GetParent()->result_id())) {
transformation_context->GetFactManager()->AddFactFunctionIsLivesafe(
message_.new_function_id());
}
if (transformation_context->GetFactManager()->BlockIsDead(
original_region_entry_block->id())) {
transformation_context->GetFactManager()->AddFactBlockIsDead(
outlined_function_ptr->entry()->id());
}
}
protobufs::Transformation TransformationOutlineFunction::ToMessage() const {
protobufs::Transformation result;
*result.mutable_outline_function() = message_;
return result;
}
std::vector<uint32_t> TransformationOutlineFunction::GetRegionInputIds(
opt::IRContext* ir_context, const std::set<opt::BasicBlock*>& region_set,
opt::BasicBlock* region_exit_block) {
std::vector<uint32_t> result;
auto enclosing_function = region_exit_block->GetParent();
enclosing_function->ForEachParam(
[ir_context, ®ion_set, &result](opt::Instruction* function_parameter) {
ir_context->get_def_use_mgr()->WhileEachUse(
function_parameter,
[ir_context, function_parameter, ®ion_set, &result](
opt::Instruction* use, uint32_t ) {
auto use_block = ir_context->get_instr_block(use);
if (use_block && region_set.count(use_block) != 0) {
result.push_back(function_parameter->result_id());
return false;
}
return true;
});
});
for (auto& block : *enclosing_function) {
std::vector<opt::Instruction*> candidate_input_ids_for_block;
if (region_set.count(&block) == 0) {
for (auto& inst : block) {
candidate_input_ids_for_block.push_back(&inst);
}
} else {
continue;
}
for (auto& inst : candidate_input_ids_for_block) {
ir_context->get_def_use_mgr()->WhileEachUse(
inst,
[ir_context, &inst, region_exit_block, ®ion_set, &result](
opt::Instruction* use, uint32_t ) -> bool {
auto use_block = ir_context->get_instr_block(use);
if (!use_block) {
return true;
}
if (region_set.count(use_block) == 0) {
return true;
}
if (use_block == region_exit_block && use->IsBlockTerminator()) {
return true;
}
result.push_back(inst->result_id());
return false;
});
}
}
return result;
}
std::vector<uint32_t> TransformationOutlineFunction::GetRegionOutputIds(
opt::IRContext* ir_context, const std::set<opt::BasicBlock*>& region_set,
opt::BasicBlock* region_exit_block) {
std::vector<uint32_t> result;
for (auto& block : *region_exit_block->GetParent()) {
if (region_set.count(&block) == 0) {
continue;
}
for (auto& inst : block) {
ir_context->get_def_use_mgr()->WhileEachUse(
&inst,
[®ion_set, ir_context, &inst, region_exit_block, &result](
opt::Instruction* use, uint32_t ) -> bool {
auto use_block = ir_context->get_instr_block(use);
if (!use_block) {
return true;
}
if (region_set.count(use_block) != 0) {
if (use_block != region_exit_block || !use->IsBlockTerminator()) {
return true;
}
}
result.push_back(inst.result_id());
return false;
});
}
}
return result;
}
std::set<opt::BasicBlock*> TransformationOutlineFunction::GetRegionBlocks(
opt::IRContext* ir_context, opt::BasicBlock* entry_block,
opt::BasicBlock* exit_block) {
auto enclosing_function = entry_block->GetParent();
auto dominator_analysis =
ir_context->GetDominatorAnalysis(enclosing_function);
auto postdominator_analysis =
ir_context->GetPostDominatorAnalysis(enclosing_function);
std::set<opt::BasicBlock*> result;
for (auto& block : *enclosing_function) {
if (dominator_analysis->Dominates(entry_block, &block) &&
postdominator_analysis->Dominates(exit_block, &block)) {
result.insert(&block);
}
}
return result;
}
std::unique_ptr<opt::Function>
TransformationOutlineFunction::PrepareFunctionPrototype(
const std::vector<uint32_t>& region_input_ids,
const std::vector<uint32_t>& region_output_ids,
const std::map<uint32_t, uint32_t>& input_id_to_fresh_id_map,
opt::IRContext* ir_context,
TransformationContext* transformation_context) const {
uint32_t return_type_id = 0;
uint32_t function_type_id = 0;
if (region_output_ids.empty()) {
std::vector<uint32_t> return_and_parameter_types;
opt::analysis::Void void_type;
return_type_id = ir_context->get_type_mgr()->GetId(&void_type);
return_and_parameter_types.push_back(return_type_id);
for (auto id : region_input_ids) {
return_and_parameter_types.push_back(
ir_context->get_def_use_mgr()->GetDef(id)->type_id());
}
function_type_id =
fuzzerutil::FindFunctionType(ir_context, return_and_parameter_types);
}
if (function_type_id == 0) {
assert(
((return_type_id == 0) == !region_output_ids.empty()) &&
"We should only have set the return type if there are no output ids.");
if (!region_output_ids.empty()) {
opt::Instruction::OperandList struct_member_types;
for (uint32_t output_id : region_output_ids) {
auto output_id_type =
ir_context->get_def_use_mgr()->GetDef(output_id)->type_id();
if (ir_context->get_def_use_mgr()->GetDef(output_id_type)->opcode() ==
spv::Op::OpTypeVoid) {
continue;
}
struct_member_types.push_back({SPV_OPERAND_TYPE_ID, {output_id_type}});
}
ir_context->module()->AddType(MakeUnique<opt::Instruction>(
ir_context, spv::Op::OpTypeStruct, 0,
message_.new_function_struct_return_type_id(),
std::move(struct_member_types)));
return_type_id = message_.new_function_struct_return_type_id();
}
assert(
return_type_id != 0 &&
"We should either have a void return type, or have created a struct.");
opt::Instruction::OperandList function_type_operands;
function_type_operands.push_back({SPV_OPERAND_TYPE_ID, {return_type_id}});
for (auto id : region_input_ids) {
function_type_operands.push_back(
{SPV_OPERAND_TYPE_ID,
{ir_context->get_def_use_mgr()->GetDef(id)->type_id()}});
}
ir_context->module()->AddType(MakeUnique<opt::Instruction>(
ir_context, spv::Op::OpTypeFunction, 0, message_.new_function_type_id(),
function_type_operands));
function_type_id = message_.new_function_type_id();
}
std::unique_ptr<opt::Function> outlined_function =
MakeUnique<opt::Function>(MakeUnique<opt::Instruction>(
ir_context, spv::Op::OpFunction, return_type_id,
message_.new_function_id(),
opt::Instruction::OperandList(
{{spv_operand_type_t ::SPV_OPERAND_TYPE_LITERAL_INTEGER,
{uint32_t(spv::FunctionControlMask::MaskNone)}},
{spv_operand_type_t::SPV_OPERAND_TYPE_ID,
{function_type_id}}})));
for (auto id : region_input_ids) {
uint32_t fresh_id = input_id_to_fresh_id_map.at(id);
outlined_function->AddParameter(MakeUnique<opt::Instruction>(
ir_context, spv::Op::OpFunctionParameter,
ir_context->get_def_use_mgr()->GetDef(id)->type_id(), fresh_id,
opt::Instruction::OperandList()));
outlined_function->ForEachParam(
[fresh_id, ir_context](opt::Instruction* inst) {
if (inst->result_id() == fresh_id) {
ir_context->AnalyzeDefUse(inst);
}
});
if (transformation_context->GetFactManager()->PointeeValueIsIrrelevant(
id)) {
transformation_context->GetFactManager()
->AddFactValueOfPointeeIsIrrelevant(input_id_to_fresh_id_map.at(id));
}
}
return outlined_function;
}
void TransformationOutlineFunction::UpdateModuleIdBoundForFreshIds(
opt::IRContext* ir_context,
const std::map<uint32_t, uint32_t>& input_id_to_fresh_id_map,
const std::map<uint32_t, uint32_t>& output_id_to_fresh_id_map) const {
fuzzerutil::UpdateModuleIdBound(
ir_context, message_.new_function_struct_return_type_id());
fuzzerutil::UpdateModuleIdBound(ir_context, message_.new_function_type_id());
fuzzerutil::UpdateModuleIdBound(ir_context, message_.new_function_id());
fuzzerutil::UpdateModuleIdBound(ir_context,
message_.new_function_region_entry_block());
fuzzerutil::UpdateModuleIdBound(ir_context, message_.new_caller_result_id());
fuzzerutil::UpdateModuleIdBound(ir_context, message_.new_callee_result_id());
for (auto& entry : input_id_to_fresh_id_map) {
fuzzerutil::UpdateModuleIdBound(ir_context, entry.second);
}
for (auto& entry : output_id_to_fresh_id_map) {
fuzzerutil::UpdateModuleIdBound(ir_context, entry.second);
}
}
void TransformationOutlineFunction::RemapInputAndOutputIdsInRegion(
opt::IRContext* ir_context,
const opt::BasicBlock& original_region_exit_block,
const std::set<opt::BasicBlock*>& region_blocks,
const std::vector<uint32_t>& region_input_ids,
const std::vector<uint32_t>& region_output_ids,
const std::map<uint32_t, uint32_t>& input_id_to_fresh_id_map,
const std::map<uint32_t, uint32_t>& output_id_to_fresh_id_map) const {
for (uint32_t id : region_input_ids) {
ir_context->get_def_use_mgr()->ForEachUse(
id, [ir_context, id, &input_id_to_fresh_id_map, region_blocks](
opt::Instruction* use, uint32_t operand_index) {
opt::BasicBlock* use_block = ir_context->get_instr_block(use);
if (region_blocks.count(use_block) != 0) {
use->SetOperand(operand_index, {input_id_to_fresh_id_map.at(id)});
}
});
}
for (uint32_t id : region_output_ids) {
ir_context->get_def_use_mgr()->ForEachUse(
id, [ir_context, &original_region_exit_block, id,
&output_id_to_fresh_id_map,
region_blocks](opt::Instruction* use, uint32_t operand_index) {
auto use_block = ir_context->get_instr_block(use);
if (
region_blocks.count(use_block) != 0 &&
!(use_block == &original_region_exit_block &&
use->IsBlockTerminator())) {
use->SetOperand(operand_index, {output_id_to_fresh_id_map.at(id)});
}
});
ir_context->get_def_use_mgr()->GetDef(id)->SetResultId(
output_id_to_fresh_id_map.at(id));
}
}
void TransformationOutlineFunction::PopulateOutlinedFunction(
const opt::BasicBlock& original_region_entry_block,
const opt::BasicBlock& original_region_exit_block,
const std::set<opt::BasicBlock*>& region_blocks,
const std::vector<uint32_t>& region_output_ids,
const std::map<uint32_t, uint32_t>& output_id_to_type_id,
const std::map<uint32_t, uint32_t>& output_id_to_fresh_id_map,
opt::IRContext* ir_context, opt::Function* outlined_function) const {
opt::BasicBlock* outlined_region_exit_block = nullptr;
std::unique_ptr<opt::BasicBlock> outlined_region_entry_block =
MakeUnique<opt::BasicBlock>(MakeUnique<opt::Instruction>(
ir_context, spv::Op::OpLabel, 0,
message_.new_function_region_entry_block(),
opt::Instruction::OperandList()));
if (&original_region_entry_block == &original_region_exit_block) {
outlined_region_exit_block = outlined_region_entry_block.get();
}
for (auto& inst : original_region_entry_block) {
outlined_region_entry_block->AddInstruction(
std::unique_ptr<opt::Instruction>(inst.Clone(ir_context)));
}
outlined_function->AddBasicBlock(std::move(outlined_region_entry_block));
auto enclosing_function = original_region_entry_block.GetParent();
for (auto block_it = enclosing_function->begin();
block_it != enclosing_function->end();) {
if (region_blocks.count(&*block_it) == 0 ||
&*block_it == &original_region_entry_block) {
++block_it;
continue;
}
auto cloned_block =
std::unique_ptr<opt::BasicBlock>(block_it->Clone(ir_context));
if (&*block_it == &original_region_exit_block) {
assert(outlined_region_exit_block == nullptr &&
"We should not yet have encountered the exit block.");
outlined_region_exit_block = cloned_block.get();
}
cloned_block->ForEachPhiInst([this](opt::Instruction* phi_inst) {
for (uint32_t predecessor_index = 1;
predecessor_index < phi_inst->NumInOperands();
predecessor_index += 2) {
if (phi_inst->GetSingleWordInOperand(predecessor_index) ==
message_.entry_block()) {
phi_inst->SetInOperand(predecessor_index,
{message_.new_function_region_entry_block()});
}
}
});
outlined_function->AddBasicBlock(std::move(cloned_block));
block_it = block_it.Erase();
}
assert(outlined_region_exit_block != nullptr &&
"We should have encountered the region's exit block when iterating "
"through the function");
for (auto inst_it = outlined_region_exit_block->begin();
inst_it != outlined_region_exit_block->end();) {
if (inst_it->opcode() == spv::Op::OpLoopMerge ||
inst_it->opcode() == spv::Op::OpSelectionMerge) {
inst_it = inst_it.Erase();
} else if (inst_it->IsBlockTerminator()) {
inst_it = inst_it.Erase();
} else {
++inst_it;
}
}
if (region_output_ids.empty()) {
outlined_region_exit_block->AddInstruction(MakeUnique<opt::Instruction>(
ir_context, spv::Op::OpReturn, 0, 0, opt::Instruction::OperandList()));
} else {
opt::Instruction::OperandList struct_member_operands;
for (uint32_t id : region_output_ids) {
if (ir_context->get_def_use_mgr()
->GetDef(output_id_to_type_id.at(id))
->opcode() != spv::Op::OpTypeVoid) {
struct_member_operands.push_back(
{SPV_OPERAND_TYPE_ID, {output_id_to_fresh_id_map.at(id)}});
}
}
outlined_region_exit_block->AddInstruction(MakeUnique<opt::Instruction>(
ir_context, spv::Op::OpCompositeConstruct,
message_.new_function_struct_return_type_id(),
message_.new_callee_result_id(), struct_member_operands));
outlined_region_exit_block->AddInstruction(MakeUnique<opt::Instruction>(
ir_context, spv::Op::OpReturnValue, 0, 0,
opt::Instruction::OperandList(
{{SPV_OPERAND_TYPE_ID, {message_.new_callee_result_id()}}})));
}
outlined_function->SetFunctionEnd(
MakeUnique<opt::Instruction>(ir_context, spv::Op::OpFunctionEnd, 0, 0,
opt::Instruction::OperandList()));
}
void TransformationOutlineFunction::ShrinkOriginalRegion(
opt::IRContext* ir_context, const std::set<opt::BasicBlock*>& region_blocks,
const std::vector<uint32_t>& region_input_ids,
const std::vector<uint32_t>& region_output_ids,
const std::map<uint32_t, uint32_t>& output_id_to_type_id,
uint32_t return_type_id,
std::unique_ptr<opt::Instruction> cloned_exit_block_merge,
std::unique_ptr<opt::Instruction> cloned_exit_block_terminator,
opt::BasicBlock* original_region_entry_block) const {
auto enclosing_function = original_region_entry_block->GetParent();
for (auto block_it = enclosing_function->begin();
block_it != enclosing_function->end();) {
if (&*block_it == original_region_entry_block) {
++block_it;
} else if (region_blocks.count(&*block_it) == 0) {
assert(block_it->MergeBlockIdIfAny() != message_.exit_block() &&
"Outlined region must not end with a merge block");
assert(block_it->ContinueBlockIdIfAny() != message_.exit_block() &&
"Outlined region must not end with a continue target");
block_it->ForEachPhiInst([this](opt::Instruction* phi_inst) {
for (uint32_t predecessor_index = 1;
predecessor_index < phi_inst->NumInOperands();
predecessor_index += 2) {
if (phi_inst->GetSingleWordInOperand(predecessor_index) ==
message_.exit_block()) {
phi_inst->SetInOperand(predecessor_index, {message_.entry_block()});
}
}
});
++block_it;
} else {
block_it = block_it.Erase();
}
}
for (auto inst_it = original_region_entry_block->begin();
inst_it != original_region_entry_block->end();) {
inst_it = inst_it.Erase();
}
opt::Instruction::OperandList function_call_operands;
function_call_operands.push_back(
{SPV_OPERAND_TYPE_ID, {message_.new_function_id()}});
for (auto input_id : region_input_ids) {
function_call_operands.push_back({SPV_OPERAND_TYPE_ID, {input_id}});
}
original_region_entry_block->AddInstruction(MakeUnique<opt::Instruction>(
ir_context, spv::Op::OpFunctionCall, return_type_id,
message_.new_caller_result_id(), function_call_operands));
uint32_t struct_member_index = 0;
for (uint32_t output_id : region_output_ids) {
uint32_t output_type_id = output_id_to_type_id.at(output_id);
if (ir_context->get_def_use_mgr()->GetDef(output_type_id)->opcode() ==
spv::Op::OpTypeVoid) {
original_region_entry_block->AddInstruction(MakeUnique<opt::Instruction>(
ir_context, spv::Op::OpUndef, output_type_id, output_id,
opt::Instruction::OperandList()));
} else {
original_region_entry_block->AddInstruction(MakeUnique<opt::Instruction>(
ir_context, spv::Op::OpCompositeExtract, output_type_id, output_id,
opt::Instruction::OperandList(
{{SPV_OPERAND_TYPE_ID, {message_.new_caller_result_id()}},
{SPV_OPERAND_TYPE_LITERAL_INTEGER, {struct_member_index}}})));
struct_member_index++;
}
}
if (cloned_exit_block_merge != nullptr) {
original_region_entry_block->AddInstruction(
std::move(cloned_exit_block_merge));
}
original_region_entry_block->AddInstruction(
std::move(cloned_exit_block_terminator));
}
std::unordered_set<uint32_t> TransformationOutlineFunction::GetFreshIds()
const {
std::unordered_set<uint32_t> result = {
message_.new_function_struct_return_type_id(),
message_.new_function_type_id(),
message_.new_function_id(),
message_.new_function_region_entry_block(),
message_.new_caller_result_id(),
message_.new_callee_result_id()};
for (auto& pair : message_.input_id_to_fresh_id()) {
result.insert(pair.second());
}
for (auto& pair : message_.output_id_to_fresh_id()) {
result.insert(pair.second());
}
return result;
}
}
}