mirror of
https://github.com/cloudwu/skynet.git
synced 2026-07-22 02:53:09 +00:00
Add async dns query
This commit is contained in:
290
lualib/dns.lua
Normal file
290
lualib/dns.lua
Normal file
@@ -0,0 +1,290 @@
|
||||
--[[
|
||||
lua dns resolver library
|
||||
See https://github.com/xjdrew/levent/blob/master/levent/dns.lua for more detail
|
||||
|
||||
-- resource record type:
|
||||
-- TYPE value and meaning
|
||||
-- A 1 a host address
|
||||
-- NS 2 an authoritative name server
|
||||
-- MD 3 a mail destination (Obsolete - use MX)
|
||||
-- MF 4 a mail forwarder (Obsolete - use MX)
|
||||
-- CNAME 5 the canonical name for an alias
|
||||
-- SOA 6 marks the start of a zone of authority
|
||||
-- MB 7 a mailbox domain name (EXPERIMENTAL)
|
||||
-- MG 8 a mail group member (EXPERIMENTAL)
|
||||
-- MR 9 a mail rename domain name (EXPERIMENTAL)
|
||||
-- NULL 10 a null RR (EXPERIMENTAL)
|
||||
-- WKS 11 a well known service description
|
||||
-- PTR 12 a domain name pointer
|
||||
-- HINFO 13 host information
|
||||
-- MINFO 14 mailbox or mail list information
|
||||
-- MX 15 mail exchange
|
||||
-- TXT 16 text strings
|
||||
-- AAAA 28 a ipv6 host address
|
||||
-- only appear in the question section:
|
||||
-- AXFR 252 A request for a transfer of an entire zone
|
||||
-- MAILB 253 A request for mailbox-related records (MB, MG or MR)
|
||||
-- MAILA 254 A request for mail agent RRs (Obsolete - see MX)
|
||||
-- * 255 A request for all records
|
||||
--
|
||||
-- resource recode class:
|
||||
-- IN 1 the Internet
|
||||
-- CS 2 the CSNET class (Obsolete - used only for examples in some obsolete RFCs)
|
||||
-- CH 3 the CHAOS class
|
||||
-- HS 4 Hesiod [Dyer 87]
|
||||
-- only appear in the question section:
|
||||
-- * 255 any class
|
||||
-- ]]
|
||||
|
||||
--[[
|
||||
-- struct header {
|
||||
-- uint16_t tid # identifier assigned by the program that generates any kind of query.
|
||||
-- uint16_t flags # flags
|
||||
-- uint16_t qdcount # the number of entries in the question section.
|
||||
-- uint16_t ancount # the number of resource records in the answer section.
|
||||
-- uint16_t nscount # the number of name server resource records in the authority records section.
|
||||
-- uint16_t arcount # the number of resource records in the additional records section.
|
||||
-- }
|
||||
--
|
||||
-- request body:
|
||||
-- struct request {
|
||||
-- string name
|
||||
-- uint16_t atype
|
||||
-- uint16_t class
|
||||
-- }
|
||||
--
|
||||
-- response body:
|
||||
-- struct response {
|
||||
-- string name
|
||||
-- uint16_t atype
|
||||
-- uint16_t class
|
||||
-- uint16_t ttl
|
||||
-- uint16_t rdlength
|
||||
-- string rdata
|
||||
-- }
|
||||
--]]
|
||||
|
||||
local skynet = require "skynet"
|
||||
local socket = require "socket"
|
||||
|
||||
local MAX_DOMAIN_LEN = 1024
|
||||
local MAX_LABEL_LEN = 63
|
||||
local MAX_PACKET_LEN = 2048
|
||||
local DNS_HEADER_LEN = 12
|
||||
|
||||
local QTYPE = {
|
||||
A = 1,
|
||||
CNAME = 5,
|
||||
AAAA = 28,
|
||||
}
|
||||
|
||||
local QCLASS = {
|
||||
IN = 1,
|
||||
}
|
||||
|
||||
local dns = {}
|
||||
|
||||
local function verify_domain_name(name)
|
||||
if #name > MAX_DOMAIN_LEN then
|
||||
return false
|
||||
end
|
||||
if not name:match("^[%l%d-%.]+$") then
|
||||
return false
|
||||
end
|
||||
for w in name:gmatch("([%w-]+)%.?") do
|
||||
if #w > MAX_LABEL_LEN then
|
||||
return false
|
||||
end
|
||||
end
|
||||
return true
|
||||
end
|
||||
|
||||
local next_tid = 1
|
||||
local function gen_tid()
|
||||
local tid = next_tid
|
||||
next_tid = next_tid + 1
|
||||
return tid
|
||||
end
|
||||
|
||||
local function pack_header(t)
|
||||
return string.pack(">HHHHHH",
|
||||
t.tid, t.flags, t.qdcount, t.ancount or 0, t.nscount or 0, t.arcount or 0)
|
||||
end
|
||||
|
||||
local function pack_question(name, qtype, qclass)
|
||||
local labels = {}
|
||||
for w in name:gmatch("([%w-]+)%.?") do
|
||||
table.insert(labels, string.pack("s1",w))
|
||||
end
|
||||
table.insert(labels, '\0')
|
||||
table.insert(labels, string.pack(">HH", qtype, qclass))
|
||||
return table.concat(labels)
|
||||
end
|
||||
|
||||
local function unpack_header(chunk)
|
||||
local tid, flags, qdcount, ancount, nscount, arcount, left = string.unpack(">HHHHHH", chunk)
|
||||
return {
|
||||
tid = tid,
|
||||
flags = flags,
|
||||
qdcount = qdcount,
|
||||
ancount = ancount,
|
||||
nscount = nscount,
|
||||
arcount = arcount
|
||||
}, left
|
||||
end
|
||||
|
||||
-- unpack a resource name
|
||||
local function unpack_name(chunk, left)
|
||||
local t = {}
|
||||
local jump_pointer
|
||||
local tag, offset, label
|
||||
while true do
|
||||
tag, left = string.unpack("B", chunk, left)
|
||||
if tag & 0xc0 == 0xc0 then
|
||||
-- pointer
|
||||
offset,left = string.unpack(">H", chunk, left - 1)
|
||||
offset = offset & 0x3fff
|
||||
if not jump_pointer then
|
||||
jump_pointer = left
|
||||
end
|
||||
-- offset is base 0, need to plus 1
|
||||
left = offset + 1
|
||||
elseif tag == 0 then
|
||||
break
|
||||
else
|
||||
label, left = string.unpack("s1", chunk, left - 1)
|
||||
t[#t+1] = label
|
||||
end
|
||||
end
|
||||
return table.concat(t, "."), jump_pointer or left
|
||||
end
|
||||
|
||||
local function unpack_question(chunk, left)
|
||||
local name, left = unpack_name(chunk, left)
|
||||
local atype, class, left = string.unpack(">HH", chunk, left)
|
||||
return {
|
||||
name = name,
|
||||
atype = atype,
|
||||
class = class
|
||||
}, left
|
||||
end
|
||||
|
||||
local function unpack_answer(chunk, left)
|
||||
local name, left = unpack_name(chunk, left)
|
||||
local atype, class, ttl, rdata, left = string.unpack(">HHI4s2", chunk, left)
|
||||
return {
|
||||
name = name,
|
||||
atype = atype,
|
||||
class = class,
|
||||
ttl = ttl,
|
||||
rdata = rdata
|
||||
},left
|
||||
end
|
||||
|
||||
local function unpack_rdata(qtype, chunk)
|
||||
if qtype == QTYPE.A then
|
||||
local a,b,c,d = string.unpack("BBBB", chunk)
|
||||
return string.format("%d.%d.%d.%d", a,b,c,d)
|
||||
elseif qtype == QTYPE.AAAA then
|
||||
local a,b,c,d,e,f,g,h = string.unpack(">HHHHHHHH", chunk)
|
||||
return string.format("%x:%x:%x:%x:%x:%x:%x:%x", a, b, c, d, e, f, g, h)
|
||||
else
|
||||
error("Error qtype " .. qtype)
|
||||
end
|
||||
end
|
||||
|
||||
local dns_server
|
||||
local request_pool = {}
|
||||
|
||||
local function resolve(content)
|
||||
if #content < DNS_HEADER_LEN then
|
||||
-- drop
|
||||
skynet.error("Recv an invalid package when dns query")
|
||||
return
|
||||
end
|
||||
local answer_header,left = unpack_header(content)
|
||||
-- verify answer
|
||||
assert(answer_header.qdcount == 1, "malformed packet")
|
||||
|
||||
local resp = request_pool[answer_header.tid]
|
||||
if not resp then
|
||||
skynet.error("Recv an invalid tid when dns query")
|
||||
return
|
||||
end
|
||||
|
||||
local question,left = unpack_question(content, left)
|
||||
if question.name ~= resp.name then
|
||||
skynet.error("Recv an invalid name when dns query")
|
||||
return
|
||||
end
|
||||
|
||||
local ttl
|
||||
local answer
|
||||
local answers = {}
|
||||
for i=1, answer_header.ancount do
|
||||
answer, left = unpack_answer(content, left)
|
||||
-- only extract qtype address
|
||||
if answer.atype == resp.qtype then
|
||||
local ip = unpack_rdata(resp.qtype, answer.rdata)
|
||||
ttl = ttl and math.min(ttl, answer.ttl) or answer.ttl
|
||||
answers[#answers+1] = ip
|
||||
end
|
||||
end
|
||||
|
||||
if #answers > 0 then
|
||||
resp.answers = answers
|
||||
end
|
||||
|
||||
skynet.wakeup(resp.co)
|
||||
end
|
||||
|
||||
function dns.server(server, port)
|
||||
if not server then
|
||||
local f = assert(io.open "/etc/resolv.conf")
|
||||
for line in f:lines() do
|
||||
server = line:match("%s*nameserver%s+([^%s]+)")
|
||||
if server then
|
||||
break
|
||||
end
|
||||
end
|
||||
assert(server, "Can't get nameserver")
|
||||
end
|
||||
dns_server = socket.udp(function(data, sz, from)
|
||||
resolve(skynet.tostring(data,sz))
|
||||
end)
|
||||
socket.udp_connect(dns_server, server, port or 53)
|
||||
return server
|
||||
end
|
||||
|
||||
local function suspend(tid, name, qtype)
|
||||
local req = {
|
||||
name = name,
|
||||
tid = tid,
|
||||
qtype = qtype,
|
||||
time = skynet.now(), -- for timeout
|
||||
co = coroutine.running(),
|
||||
}
|
||||
request_pool[tid] = req
|
||||
skynet.wait()
|
||||
local answers = request_pool[tid].answers
|
||||
request_pool[tid] = nil
|
||||
assert(answers, "no ip")
|
||||
return answers[1], answers
|
||||
end
|
||||
|
||||
function dns.resolve(name, ipv6)
|
||||
local qtype = ipv6 and QTYPE.AAAA or QTYPE.A
|
||||
local name = name:lower()
|
||||
assert(verify_domain_name(name) , "illegal name")
|
||||
local question_header = {
|
||||
tid = gen_tid(),
|
||||
flags = 0x100, -- flags: 00000001 00000000, set RD
|
||||
qdcount = 1,
|
||||
}
|
||||
local req = pack_header(question_header) .. pack_question(name, qtype, QCLASS.IN)
|
||||
assert(dns_server, "Call dns.server fist")
|
||||
socket.write(dns_server, req)
|
||||
return suspend(question_header.tid, name, qtype)
|
||||
end
|
||||
|
||||
return dns
|
||||
11
test/testdns.lua
Normal file
11
test/testdns.lua
Normal file
@@ -0,0 +1,11 @@
|
||||
local skynet = require "skynet"
|
||||
local dns = require "dns"
|
||||
|
||||
skynet.start(function()
|
||||
print("nameserver:", dns.server()) -- set nameserver
|
||||
-- you can specify the server like dns.server("8.8.4.4", 53)
|
||||
local ip, ips = dns.resolve "github.com"
|
||||
for k,v in ipairs(ips) do
|
||||
print("github.com",v)
|
||||
end
|
||||
end)
|
||||
Reference in New Issue
Block a user