#include "sse2.h"
typedef MASKTYPE tp_mask;
typedef MBTYPE tp_vector;
inline tp_mask kill_first(long x){	return ((tp_mask)(-1))<<x;}
inline tp_mask kill_last (long x){	return ((tp_mask)(-1))>>x;}

inline tp_vector zerov(){
	tp_vector m;
	m=XOR(m,m);
}
inline long first_bit(tp_mask t){ return __builtin_ctzl(t);}
#define find_zeros(x) get_mask(test_eq(x,vzero))
char * cmp_z(char *a,char *b,int dir){
	while(*a && *a!=*b) {a+=dir;b+=dir;}
	return a;
}

char *strstr2(char *s,char *n){int i;
	const int scan=2,unroll=4;
	tp_mask mask,zmask;
	tp_vector vzero=zerov();
	tp_vector m[scan];
	for(i=0;i<scan;i++) m[i]=make_mask(n[i],0);
	int offset=((long)s)%(unroll*BYTES_AT_ONCE);
	char *s2=s-offset;
	tp_mask mkill=kill_first(offset);
	tp_mask zkill=kill_first(offset);
	tp_vector so,sn; sn=LOAD(s2);	
	while(1){
		mask=0;
		zmask=0;
		for(i=0;i<unroll;i++){
			zmask=zmask|(find_zeros(sn)<<(i*BYTES_AT_ONCE));
			so=sn; 
			if(i==unroll-1 && (zmask&zkill)){
				int last=first_bit(zmask);
				for (i=0;i<last;i++) 
					if (!*cmp_z(n,s2+i,1)) return s2+i;
				return NULL;
			} 
			sn=LOAD(s2+(i+1)*BYTES_AT_ONCE);
			tp_vector e;/*ssse3 has no instruction for variable length shift*/
			if(scan-1>=0){	e=        XOR(CONCAT(so,sn,0),m[0]) ;}
			if(scan-1>=1){  e=  OR(e, XOR(CONCAT(so,sn,1),m[1]));}
			if(scan-1>=2){  e=  OR(e, XOR(CONCAT(so,sn,2),m[2]));}
			if(scan-1>=3){  e=  OR(e, XOR(CONCAT(so,sn,3),m[3]));}

			mask=mask|(find_zeros(e)<<(i*BYTES_AT_ONCE));
		}
		mask=mask&mkill;
		if(mask){
			while(mask){
				int i=first_bit(mask);
				char *p=s2+i;
				if (!*cmp_z(p+scan,n+scan,1)) return p;
				mkill=kill_first(i);
				mask=mask&mkill;
			}
		}
		mkill=0;
		zkill=0;
	}
}
