//
//  trie.hpp
//  aho-corasick
//
//  Created by chenguanglv on 15/12/19.
//  Copyright © 2015年 ibrain. All rights reserved.
//

#ifndef trie_h
#define trie_h

#include "typedef.hpp"
#include "state.hpp"
#include "emit.hpp"
#include "utf8_len_table.hpp"

namespace ac
{
    template<typename T,typename D>
    class Trie
    {
        typedef State<T,D> StateT;
        typedef State<T,D>* StatePtrT;
        typedef std::queue<StatePtrT> StateQueue;
        typedef std::unique_ptr<StateT> StatePtrTS;
        typedef EmitResult<T, D> EmitResultT;
        typedef typename ACDef<T>::ACData TrieData;
    public:
        typedef std::vector<EmitResultT> EmitResultsT;
        
    public:
        Trie() : _root(new StateT()){};
        
        ~Trie(){};
        
        void insert(TrieData data,D *userdata = nullptr)
        {
            if (data.size() <= 0)
            {
                return ;
            }
            
            StatePtrT state = _root.get();
            for(const auto& d : data)
            {
                state = state->addState(d);
            }
            state->addEmit(data,userdata);
            
            _needRebuild = true;
        }
        
        void build()
        {
            StateQueue q;
            for (auto& state : _root->getStates())
            {
                state->setFailure(_root.get());
                q.push(state);
            }
            
            while (!q.empty())
            {
                auto state = q.front();
                for (const auto& var : state->getVariables())
                {
                    StatePtrT nextState = state->nextState(var);
                    q.push(nextState);
                    
                    StatePtrT failureState = state->failure();
                    while(failureState->nextState(var) == nullptr)
                    {
                        failureState = failureState->failure();
                    }
                    
                    StatePtrT newFailureState = failureState->nextState(var);
                    nextState->setFailure(newFailureState);
                    nextState->addEmit(newFailureState->getEmits());
                }
                
                q.pop();
            }
            
            _needRebuild = false;
        }
        
        EmitResultsT parse(TrieData data)
        {
            if (_needRebuild)
            {
                build();
            }
            
            size_t pos = 0;
            
#ifdef USE_UTF8_TABLE
            size_t charactorPos = 0;
            unsigned char utf8ByteCount = 0;
#endif
            StatePtrT state = _root.get();
            EmitResultsT emits;
            
            for (auto& d : data)
            {
#ifdef USE_UTF8_TABLE
                state = getState(state, d);
                if(utf8ByteCount <= 0)
                {
                    utf8ByteCount = UTF8_BYTE_LEN((unsigned char)d);
                }

                utf8ByteCount--;
                
                if(utf8ByteCount == 0)
                {
                    charactorPos++;
                    storeEmits(charactorPos,state,emits);
                    
                }
                pos++;
#else
                state = getState(state, d);
                storeEmits(pos,state,emits);
                pos++;
                
#endif
            }
            return emits;
        }
        
    private:
        StatePtrT getState(StatePtrT state,T d) const
        {
            StatePtrT nextState = state->nextState(d);
            StatePtrT tempState = state;
            while (!nextState)
            {
                tempState = tempState->failure();
                nextState = tempState->nextState(d);
            }
            
            return nextState;
        }

        void storeEmits(size_t pos,StatePtrT state,EmitResultsT& emits) const
        {
            auto stateEmits = state->getEmits();
            if (!stateEmits.empty())
            {
                for (const auto& data : stateEmits)
                {
                    emits.push_back(EmitResultT(
#ifdef USE_UTF8_TABLE
                                                pos - data.getCharacterCount(),
#else
                                                pos - data.getEmitResult().size() + 1,
#endif
                                                pos,
                                                data.getEmitResult(),
                                                data.getUserData()
                                                ));
                }
            }
        }
        
        
    private:
        StatePtrTS _root;
        bool _needRebuild;
    };
}

#endif /* trie_h */
