package dns import "fmt" type unpacker struct { buf []byte off int } func (u *unpacker) readUint16() (uint16, error) { if u.off+2 > len(u.buf) { return 0, fmt.Errorf("dns: short uint16 at %d", u.off) } v := uint16(u.buf[u.off])<<8 | uint16(u.buf[u.off+1]) u.off += 2 return v, nil } func (u *unpacker) readUint32() (uint32, error) { if u.off+4 > len(u.buf) { return 0, fmt.Errorf("dns: short uint32 at %d", u.off) } v := uint32(u.buf[u.off])<<24 | uint32(u.buf[u.off+1])<<16 | uint32(u.buf[u.off+2])<<8 | uint32(u.buf[u.off+3]) u.off += 4 return v, nil } func (u *unpacker) readName() (Name, error) { var name Name start := u.off jumped := false jumps := 0 for { if u.off >= len(u.buf) { return Name{}, fmt.Errorf("dns: name overflow") } c := u.buf[u.off] // Compression pointer if c&0xC0 == 0xC0 { if u.off+2 > len(u.buf) { return Name{}, fmt.Errorf("dns: short pointer") } ptr := int(uint16(c&0x3F)<<8 | uint16(u.buf[u.off+1])) if ptr >= len(u.buf) { return Name{}, fmt.Errorf("dns: pointer out of range") } if !jumped { start = u.off + 2 } u.off = ptr jumped = true jumps++ if jumps > 10 { return Name{}, fmt.Errorf("dns: too many compression jumps") } continue } // Reserved label type (0x40, 0x80) if c&0xC0 != 0 { return Name{}, fmt.Errorf("dns: reserved label type") } // Root label if c == 0 { name.Data = append(name.Data, 0) u.off++ break } // Normal label length := int(c) if length > 63 { return Name{}, fmt.Errorf("dns: label too long") } if u.off+1+length > len(u.buf) { return Name{}, fmt.Errorf("dns: short label") } name.Data = append(name.Data, u.buf[u.off:u.off+1+length]...) u.off += 1 + length } if jumped { u.off = start } return name, nil } func Unpack(buf []byte) (*Msg, error) { if len(buf) < 12 { return nil, fmt.Errorf("dns: message too short") } u := &unpacker{buf: buf} m := &Msg{} m.Header.ID, _ = u.readUint16() m.Header.Flags, _ = u.readUint16() qdcount, _ := u.readUint16() ancount, _ := u.readUint16() nscount, _ := u.readUint16() arcount, _ := u.readUint16() m.Question = make([]Question, qdcount) for i := range m.Question { if err := m.Question[i].unpack(u); err != nil { return nil, err } } m.Answer = make([]RR, ancount) for i := range m.Answer { rr, err := unpackRR(u) if err != nil { return nil, err } m.Answer[i] = rr } m.Ns = make([]RR, nscount) for i := range m.Ns { rr, err := unpackRR(u) if err != nil { return nil, err } m.Ns[i] = rr } m.Extra = make([]RR, arcount) for i := range m.Extra { rr, err := unpackRR(u) if err != nil { return nil, err } m.Extra[i] = rr } return m, nil }