[ Web Proxy ]
URL:
Viewing: https://raw.githubusercontent.com/rburkholder/daScript/master/src/ast/ast_block_folding.cpp [Back]  [Original]

#include "daScript/misc/platform.h"

#include "daScript/ast/ast.h"
#include "daScript/ast/ast_visitor.h"

namespace das {

    class UnsafeFolding : public PassVisitor {
    protected:
        virtual ExpressionPtr visit ( ExprUnsafe * expr ) {
            return expr->body;
        }
    };

    void Program::foldUnsafe() {
        UnsafeFolding context;
        visit(context);
    }


    // this folds the following, by setting r2v flag on expressions
    //  r2v(var)            = @var
    //  r2v(expr.field)     = expr.@field
    //  r2v(expr[index])    = expr@[index]
    //  r2v(a ? b : c)      = a ? r2v(b) : r2v(c)
    //  r2v(cast(x))        = cast(r2v(x))
    class RefFolding : public PassVisitor {
    protected:
        virtual ExpressionPtr visit ( ExprRef2Value * expr ) override {
            if (expr->type->baseType == Type::tHandle) {
                return Visitor::visit(expr);
            }
            if ( expr->subexpr->rtti_isCast() ) {
                reportFolding();
                auto ecast = static_pointer_cast(expr->subexpr);
                auto nr2v = make_smart();
                nr2v->at = expr->at;
                nr2v->subexpr = ecast->subexpr;
                nr2v->type = make_smart(*nr2v->subexpr->type);
                nr2v->type->ref = false;
                ecast->subexpr = nr2v;
                ecast->type->ref = false;
                return ecast;
            } else if ( expr->subexpr->rtti_isVar() ) {
                if ( expr->subexpr->type->isHandle() ) {
                    return Visitor::visit(expr);
                } else {
                    reportFolding();
                    auto evar = static_pointer_cast(expr->subexpr);
                    evar->r2v = true;
                    evar->type->ref = false;
                    return evar;
                }
            } else if ( expr->subexpr->rtti_isField() ) {
                reportFolding();
                auto efield = static_pointer_cast(expr->subexpr);
                efield->r2v = true;
                efield->type->ref = false;
                return efield;
            } else if ( expr->subexpr->rtti_isAsVariant() ) {
                reportFolding();
                auto efield = static_pointer_cast(expr->subexpr);
                efield->r2v = true;
                efield->type->ref = false;
                return efield;
            } else if ( expr->subexpr->rtti_isSafeAsVariant() ) {
                reportFolding();
                auto efield = static_pointer_cast(expr->subexpr);
                efield->r2v = true;
                efield->type->ref = false;
                return efield;
            } else if ( expr->subexpr->rtti_isSwizzle() ) {
                reportFolding();
                auto eswiz = static_pointer_cast(expr->subexpr);
                eswiz->r2v = true;
                eswiz->type->ref = false;
                if (!TypeDecl::isSequencialMask(eswiz->fields)) {
                    eswiz->value = Expression::autoDereference(eswiz->value);
                }
                return eswiz;
            } else if ( expr->subexpr->rtti_isSafeField() ) {
                DAS_ASSERTF(false, "we should not be here. R2V of ?. is strange indeed");
                reportFolding();
                auto efield = static_pointer_cast(expr->subexpr);
                efield->r2v = true;
                efield->type->ref = false;
                return efield;
            } else if ( expr->subexpr->rtti_isAt() ) {
                reportFolding();
                auto eat = static_pointer_cast(expr->subexpr);
                eat->r2v = true;
                eat->type->ref = false;
                return eat;
            } else if ( expr->subexpr->rtti_isSafeAt() ) {
                DAS_ASSERTF(false, "we should not be here. R2V of ?[ is strange indeed");
                reportFolding();
                auto eat = static_pointer_cast(expr->subexpr);
                eat->r2v = true;
                eat->type->ref = false;
                return eat;
            } else if ( expr->subexpr->rtti_isOp3() ) {
                reportFolding();
                auto op3 = static_pointer_cast(expr->subexpr);
                op3->left = Expression::autoDereference(op3->left);
                op3->right = Expression::autoDereference(op3->right);
                op3->type->ref = false;
                return expr->subexpr;
            } else if ( expr->subexpr->rtti_isNullCoalescing() ) {
                reportFolding();
                auto nc = static_pointer_cast(expr->subexpr);
                nc->defaultValue = Expression::autoDereference(nc->defaultValue);
                nc->type->ref = false;
                return nc;
            } else {
                return Visitor::visit(expr);
            }
        }
    };

    class BlockFolding : public PassVisitor {
    protected:
        das_set labels;
        bool allLabels = false;
    protected:
        bool hasLabels ( vector & blockList ) const {
            for ( auto & expr : blockList ) {
                if ( !expr ) continue;
                if ( expr->rtti_isLabel() ) {
                    return true;
                } else  if ( expr->rtti_isBlock() ) {
                    auto pBlock = static_pointer_cast(expr);
                    if ( !pBlock->isClosure ) {
                        if ( hasLabels(pBlock->list) ) return true;
                    }
                }
            }
            return false;
        }
        void collect ( vector & list, vector & blockList ) {
            bool stopAtExit = !hasLabels(blockList);
            bool skipTilLabel = false;
            for ( auto & expr : blockList ) {
                if ( !expr ) continue;
                if ( expr->rtti_isLabel() ) {
                    auto lexpr = static_pointer_cast(expr);
                    if ( allLabels || (labels.find(lexpr->label)!=labels.end()) ) {
                        list.push_back(expr);
                        skipTilLabel = false;
                    }
                    continue;
                }
                if ( skipTilLabel ) continue;
                if ( expr->rtti_isGoto() ) {
                    list.push_back(expr);
                    skipTilLabel = true;
                    continue;
                }
                if ( expr->rtti_isBreak() || expr->rtti_isReturn() || expr->rtti_isContinue() ) {
                    if ( stopAtExit ) {
                        list.push_back(expr);
                        break;
                    } else {
                        list.push_back(expr);
                        skipTilLabel = true;
                        continue;
                    }
                }
                if ( expr->rtti_isBlock() ) {
                    auto pBlock = static_pointer_cast(expr);
                    if ( !pBlock->isClosure && !pBlock->finalList.size() ) {
                        collect(list, pBlock->list);
                    } else {
                        list.push_back(expr);
                    }
                } else if ( expr->rtti_isWith() ) {
                    auto pWith = static_pointer_cast(expr);
                    if ( pWith->body ) {
                        list.push_back(pWith->body);
                    }
                } else if ( expr->rtti_isAssume() ) {
                    // do nothing with assume
                } else {
                    list.push_back(expr);
                }
            }
        }
    protected:
        virtual void preVisit ( ExprGoto * expr ) override {
            Visitor::preVisit(expr);
            if ( expr->subexpr ) {
                allLabels = true;
            } else {
                labels.insert(expr->label);
            }
        }
    // ExprBlock
        virtual ExpressionPtr visit ( ExprBlock * block ) override {
            vector list;
            collect(list, block->list);
            if ( list!=block->list ) {
                swap ( block->list, list );
                reportFolding();
            }
            vector finalList;
            collect(finalList, block->finalList);
            if ( finalList!=block->finalList ) {
                swap ( block->finalList, finalList );
                reportFolding();
            }
            return Visitor::visit(block);
        }
    // ExprLet
        virtual ExpressionPtr visit ( ExprLet * let ) override {
            if ( let->variables.size()==0 ) {
                reportFolding();
                return nullptr;
            }
            return Visitor::visit(let);
        }
    // function
        virtual FunctionPtr visit ( Function * func ) override {
            labels.clear();
            allLabels = false;
            if ( func->body && func->result->isVoid() ) {   // remove trailing return on the void function
                if ( func->body->rtti_isBlock() ) {
                    auto block = static_pointer_cast(func->body);
                    if ( block->list.size() && block->list.back()->rtti_isReturn() ) {
                        block->list.resize(block->list.size()-1);
                        reportFolding();
                        return Visitor::visit(func);
                    }
                }
            }
            return Visitor::visit(func);
        }
    };

    class CondFolding : public PassVisitor {
    protected:
        Function * func = nullptr;
        virtual void preVisit ( Function * f ) override {
            Visitor::preVisit(f);
            func = f;
        }
        virtual FunctionPtr visit ( Function * f ) override {
            func = nullptr;
            return Visitor::visit(f);
        }
        virtual ExpressionPtr visit ( ExprIfThenElse * expr ) override {
            // if ( func && func->generator ) return Visitor::visit(expr);
            // if (cond) return x; else return y; => (cond ? x : y)
            if (expr->if_false) {
                smart_ptr lr, rr;
                if (expr->if_true->rtti_isBlock()) {
                    auto tb = static_pointer_cast(expr->if_true);
                    if (tb->list.size() == 1 && tb->list.back()->rtti_isReturn()) {
                        lr = static_pointer_cast(tb->list.back());
                        if ( lr->subexpr && lr->subexpr->rtti_isMakeLocal() ) {
                            lr.reset();   // we don't touch CMRES stuff
                        }
                    }
                }
                if (expr->if_false->rtti_isBlock()) {
                    auto fb = static_pointer_cast(expr->if_false);
                    if (fb->list.size() == 1 && fb->list.back()->rtti_isReturn()) {
                        rr = static_pointer_cast(fb->list.back());
                        if ( rr->subexpr && rr->subexpr->rtti_isMakeLocal() ) {
                            rr.reset();   // we don't touch CMRES stuff
                        }
                    }
                }
                if (lr && rr) {
                    if ( lr->moveSemantics != rr->moveSemantics ) {
                        lr.reset(); // move semantics must match
                        rr.reset();
                    }
                }
                if (lr && rr) {
                    if ( lr->subexpr ) {
                        auto newCond = make_smart(expr->at, "?", expr->cond, lr->subexpr, rr->subexpr);
                        newCond->type = make_smart(*lr->subexpr->type);
                        auto newRet = make_smart(expr->at, newCond);
                        newRet->moveSemantics = lr->moveSemantics;
                        reportFolding();
                        return newRet;
                    } else {
                        // this is actually if ( a ) return; else return;
                        reportFolding();
                        return lr;
                    }
                }
            }
            return Visitor::visit(expr);
        }
        // ExprBlock
        virtual ExpressionPtr visit ( ExprBlock * block ) override {
            if ( func && func->generator ) return Visitor::visit(block);
            /*
            if ( cond )
                ...
                break or return or continue
            b
                =>
            if ( cond )
                ...
                break or return or continue
            else
                b
            */
            if (!block->isClosure && block->list.size() > 1) {
                for ( int i=0, is=int(block->list.size())-1; i!=is; ++i ) {
                    auto expr = block->list[i];
                    if (expr != block->list.back()) {
                        if (expr->rtti_isIfThenElse()) {
                            auto ite = static_pointer_cast(expr);
                            if (!ite->if_false) {
                                if (ite->if_true->rtti_isBlock()) {
                                    auto tb = static_pointer_cast(ite->if_true);
                                    if ( tb->list.size() ) {
                                        auto lastE = tb->list.back();
                                        if (lastE->rtti_isReturn() || lastE->rtti_isBreak() || lastE->rtti_isContinue()) {
                                            vector tail;
                                            tail.insert(tail.begin(), block->list.begin() + i + 1, block->list.end());
                                            auto fb = make_smart();
                                            fb->at = tail.front()->at;
                                            swap(fb->list, tail);
                                            ite->if_false = fb;
                                            block->list.resize(i + 1);
                                            reportFolding();
                                            return Visitor::visit(block);
                                        }
                                    }
                                }
                            }
                        }
                    }
                }
            }
            return Visitor::visit(block);
        }
    };

    // program

    bool Program::optimizationRefFolding() {
        bool any = false, anything = false;
        do {
            RefFolding context;
            visit(context);
            any = context.didAnything();
            anything |= any;
        } while ( any );
        return anything;
    }

    bool Program::optimizationBlockFolding() {
        BlockFolding context;
        visit(context);
        return context.didAnything();
    }

    bool Program::optimizationCondFolding() {
        CondFolding context;
        visit(context);
        return context.didAnything();
    }

}


Web Proxy Viewer  |  New URL  |  Original Page