17.12
This commit is contained in:
@@ -0,0 +1,8 @@
|
||||
# -*- test-case-name: twisted.words.test -*-
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
Twisted Words: Client and server implementations for IRC, XMPP, and other chat
|
||||
services.
|
||||
"""
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,293 @@
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
IRC support for Instance Messenger.
|
||||
"""
|
||||
|
||||
from twisted.words.protocols import irc
|
||||
from twisted.words.im.locals import ONLINE
|
||||
from twisted.internet import defer, reactor, protocol
|
||||
from twisted.internet.defer import succeed
|
||||
from twisted.words.im import basesupport, interfaces, locals
|
||||
from zope.interface import implementer
|
||||
|
||||
|
||||
class IRCPerson(basesupport.AbstractPerson):
|
||||
|
||||
def imperson_whois(self):
|
||||
if self.account.client is None:
|
||||
raise locals.OfflineError
|
||||
self.account.client.sendLine("WHOIS %s" % self.name)
|
||||
|
||||
|
||||
### interface impl
|
||||
def isOnline(self):
|
||||
return ONLINE
|
||||
|
||||
|
||||
def getStatus(self):
|
||||
return ONLINE
|
||||
|
||||
|
||||
def setStatus(self,status):
|
||||
self.status=status
|
||||
self.chat.getContactsList().setContactStatus(self)
|
||||
|
||||
|
||||
def sendMessage(self, text, meta=None):
|
||||
if self.account.client is None:
|
||||
raise locals.OfflineError
|
||||
for line in text.split('\n'):
|
||||
if meta and meta.get("style", None) == "emote":
|
||||
self.account.client.ctcpMakeQuery(self.name,[('ACTION', line)])
|
||||
else:
|
||||
self.account.client.msg(self.name, line)
|
||||
return succeed(text)
|
||||
|
||||
|
||||
|
||||
@implementer(interfaces.IGroup)
|
||||
class IRCGroup(basesupport.AbstractGroup):
|
||||
def imgroup_testAction(self):
|
||||
pass
|
||||
|
||||
|
||||
def imtarget_kick(self, target):
|
||||
if self.account.client is None:
|
||||
raise locals.OfflineError
|
||||
reason = "for great justice!"
|
||||
self.account.client.sendLine("KICK #%s %s :%s" % (
|
||||
self.name, target.name, reason))
|
||||
|
||||
|
||||
### Interface Implementation
|
||||
def setTopic(self, topic):
|
||||
if self.account.client is None:
|
||||
raise locals.OfflineError
|
||||
self.account.client.topic(self.name, topic)
|
||||
|
||||
|
||||
def sendGroupMessage(self, text, meta={}):
|
||||
if self.account.client is None:
|
||||
raise locals.OfflineError
|
||||
if meta and meta.get("style", None) == "emote":
|
||||
self.account.client.ctcpMakeQuery(self.name,[('ACTION', text)])
|
||||
return succeed(text)
|
||||
#standard shmandard, clients don't support plain escaped newlines!
|
||||
for line in text.split('\n'):
|
||||
self.account.client.say(self.name, line)
|
||||
return succeed(text)
|
||||
|
||||
|
||||
def leave(self):
|
||||
if self.account.client is None:
|
||||
raise locals.OfflineError
|
||||
self.account.client.leave(self.name)
|
||||
self.account.client.getGroupConversation(self.name,1)
|
||||
|
||||
|
||||
|
||||
class IRCProto(basesupport.AbstractClientMixin, irc.IRCClient):
|
||||
def __init__(self, account, chatui, logonDeferred=None):
|
||||
basesupport.AbstractClientMixin.__init__(self, account, chatui,
|
||||
logonDeferred)
|
||||
self._namreplies={}
|
||||
self._ingroups={}
|
||||
self._groups={}
|
||||
self._topics={}
|
||||
|
||||
|
||||
def getGroupConversation(self, name, hide=0):
|
||||
name = name.lower()
|
||||
return self.chat.getGroupConversation(self.chat.getGroup(name, self),
|
||||
stayHidden=hide)
|
||||
|
||||
|
||||
def getPerson(self,name):
|
||||
return self.chat.getPerson(name, self)
|
||||
|
||||
|
||||
def connectionMade(self):
|
||||
# XXX: Why do I duplicate code in IRCClient.register?
|
||||
try:
|
||||
self.performLogin = True
|
||||
self.nickname = self.account.username
|
||||
self.password = self.account.password
|
||||
self.realname = "Twisted-IM user"
|
||||
|
||||
irc.IRCClient.connectionMade(self)
|
||||
|
||||
for channel in self.account.channels:
|
||||
self.joinGroup(channel)
|
||||
self.account._isOnline=1
|
||||
if self._logonDeferred is not None:
|
||||
self._logonDeferred.callback(self)
|
||||
self.chat.getContactsList()
|
||||
except:
|
||||
import traceback
|
||||
traceback.print_exc()
|
||||
|
||||
|
||||
def setNick(self,nick):
|
||||
self.name=nick
|
||||
self.accountName="%s (IRC)"%nick
|
||||
irc.IRCClient.setNick(self,nick)
|
||||
|
||||
|
||||
def kickedFrom(self, channel, kicker, message):
|
||||
"""
|
||||
Called when I am kicked from a channel.
|
||||
"""
|
||||
return self.chat.getGroupConversation(
|
||||
self.chat.getGroup(channel[1:], self), 1)
|
||||
|
||||
|
||||
def userKicked(self, kickee, channel, kicker, message):
|
||||
pass
|
||||
|
||||
|
||||
def noticed(self, username, channel, message):
|
||||
self.privmsg(username, channel, message, {"dontAutoRespond": 1})
|
||||
|
||||
|
||||
def privmsg(self, username, channel, message, metadata=None):
|
||||
if metadata is None:
|
||||
metadata = {}
|
||||
username = username.split('!',1)[0]
|
||||
if username==self.name: return
|
||||
if channel[0]=='#':
|
||||
group=channel[1:]
|
||||
self.getGroupConversation(group).showGroupMessage(username, message, metadata)
|
||||
return
|
||||
self.chat.getConversation(self.getPerson(username)).showMessage(message, metadata)
|
||||
|
||||
|
||||
def action(self,username,channel,emote):
|
||||
username = username.split('!',1)[0]
|
||||
if username==self.name: return
|
||||
meta={'style':'emote'}
|
||||
if channel[0]=='#':
|
||||
group=channel[1:]
|
||||
self.getGroupConversation(group).showGroupMessage(username, emote, meta)
|
||||
return
|
||||
self.chat.getConversation(self.getPerson(username)).showMessage(emote,meta)
|
||||
|
||||
|
||||
def irc_RPL_NAMREPLY(self,prefix,params):
|
||||
"""
|
||||
RPL_NAMREPLY
|
||||
>> NAMES #bnl
|
||||
<< :Arlington.VA.US.Undernet.Org 353 z3p = #bnl :pSwede Dan-- SkOyg AG
|
||||
"""
|
||||
group = params[2][1:].lower()
|
||||
users = params[3].split()
|
||||
for ui in range(len(users)):
|
||||
while users[ui][0] in ["@","+"]: # channel modes
|
||||
users[ui]=users[ui][1:]
|
||||
if group not in self._namreplies:
|
||||
self._namreplies[group]=[]
|
||||
self._namreplies[group].extend(users)
|
||||
for nickname in users:
|
||||
try:
|
||||
self._ingroups[nickname].append(group)
|
||||
except:
|
||||
self._ingroups[nickname]=[group]
|
||||
|
||||
|
||||
def irc_RPL_ENDOFNAMES(self,prefix,params):
|
||||
group=params[1][1:]
|
||||
self.getGroupConversation(group).setGroupMembers(self._namreplies[group.lower()])
|
||||
del self._namreplies[group.lower()]
|
||||
|
||||
|
||||
def irc_RPL_TOPIC(self,prefix,params):
|
||||
self._topics[params[1][1:]]=params[2]
|
||||
|
||||
|
||||
def irc_333(self,prefix,params):
|
||||
group=params[1][1:]
|
||||
self.getGroupConversation(group).setTopic(self._topics[group],params[2])
|
||||
del self._topics[group]
|
||||
|
||||
|
||||
def irc_TOPIC(self,prefix,params):
|
||||
nickname = prefix.split("!")[0]
|
||||
group = params[0][1:]
|
||||
topic = params[1]
|
||||
self.getGroupConversation(group).setTopic(topic,nickname)
|
||||
|
||||
|
||||
def irc_JOIN(self,prefix,params):
|
||||
nickname = prefix.split("!")[0]
|
||||
group = params[0][1:].lower()
|
||||
if nickname!=self.nickname:
|
||||
try:
|
||||
self._ingroups[nickname].append(group)
|
||||
except:
|
||||
self._ingroups[nickname]=[group]
|
||||
self.getGroupConversation(group).memberJoined(nickname)
|
||||
|
||||
|
||||
def irc_PART(self,prefix,params):
|
||||
nickname = prefix.split("!")[0]
|
||||
group = params[0][1:].lower()
|
||||
if nickname!=self.nickname:
|
||||
if group in self._ingroups[nickname]:
|
||||
self._ingroups[nickname].remove(group)
|
||||
self.getGroupConversation(group).memberLeft(nickname)
|
||||
|
||||
|
||||
def irc_QUIT(self,prefix,params):
|
||||
nickname = prefix.split("!")[0]
|
||||
if nickname in self._ingroups:
|
||||
for group in self._ingroups[nickname]:
|
||||
self.getGroupConversation(group).memberLeft(nickname)
|
||||
self._ingroups[nickname]=[]
|
||||
|
||||
|
||||
def irc_NICK(self, prefix, params):
|
||||
fromNick = prefix.split("!")[0]
|
||||
toNick = params[0]
|
||||
if fromNick not in self._ingroups:
|
||||
return
|
||||
for group in self._ingroups[fromNick]:
|
||||
self.getGroupConversation(group).memberChangedNick(fromNick, toNick)
|
||||
self._ingroups[toNick] = self._ingroups[fromNick]
|
||||
del self._ingroups[fromNick]
|
||||
|
||||
|
||||
def irc_unknown(self, prefix, command, params):
|
||||
pass
|
||||
|
||||
|
||||
# GTKIM calls
|
||||
def joinGroup(self,name):
|
||||
self.join(name)
|
||||
self.getGroupConversation(name)
|
||||
|
||||
|
||||
|
||||
@implementer(interfaces.IAccount)
|
||||
class IRCAccount(basesupport.AbstractAccount):
|
||||
gatewayType = "IRC"
|
||||
|
||||
_groupFactory = IRCGroup
|
||||
_personFactory = IRCPerson
|
||||
|
||||
def __init__(self, accountName, autoLogin, username, password, host, port,
|
||||
channels=''):
|
||||
basesupport.AbstractAccount.__init__(self, accountName, autoLogin,
|
||||
username, password, host, port)
|
||||
self.channels = [channel.strip() for channel in channels.split(',')]
|
||||
if self.channels == ['']:
|
||||
self.channels = []
|
||||
|
||||
|
||||
def _startLogOn(self, chatui):
|
||||
logonDeferred = defer.Deferred()
|
||||
cc = protocol.ClientCreator(reactor, IRCProto, self, chatui,
|
||||
logonDeferred)
|
||||
d = cc.connectTCP(self.host, self.port)
|
||||
d.addErrback(logonDeferred.errback)
|
||||
return logonDeferred
|
||||
@@ -0,0 +1,262 @@
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
|
||||
"""
|
||||
L{twisted.words} support for Instance Messenger.
|
||||
"""
|
||||
|
||||
from __future__ import print_function
|
||||
|
||||
from twisted.internet import defer
|
||||
from twisted.internet import error
|
||||
from twisted.python import log
|
||||
from twisted.python.failure import Failure
|
||||
from twisted.spread import pb
|
||||
|
||||
from twisted.words.im.locals import ONLINE, OFFLINE, AWAY
|
||||
|
||||
from twisted.words.im import basesupport, interfaces
|
||||
from zope.interface import implementer
|
||||
|
||||
|
||||
class TwistedWordsPerson(basesupport.AbstractPerson):
|
||||
"""I a facade for a person you can talk to through a twisted.words service.
|
||||
"""
|
||||
def __init__(self, name, wordsAccount):
|
||||
basesupport.AbstractPerson.__init__(self, name, wordsAccount)
|
||||
self.status = OFFLINE
|
||||
|
||||
def isOnline(self):
|
||||
return ((self.status == ONLINE) or
|
||||
(self.status == AWAY))
|
||||
|
||||
def getStatus(self):
|
||||
return self.status
|
||||
|
||||
def sendMessage(self, text, metadata):
|
||||
"""Return a deferred...
|
||||
"""
|
||||
if metadata:
|
||||
d=self.account.client.perspective.directMessage(self.name,
|
||||
text, metadata)
|
||||
d.addErrback(self.metadataFailed, "* "+text)
|
||||
return d
|
||||
else:
|
||||
return self.account.client.perspective.callRemote('directMessage',self.name, text)
|
||||
|
||||
def metadataFailed(self, result, text):
|
||||
print("result:",result,"text:",text)
|
||||
return self.account.client.perspective.directMessage(self.name, text)
|
||||
|
||||
def setStatus(self, status):
|
||||
self.status = status
|
||||
self.chat.getContactsList().setContactStatus(self)
|
||||
|
||||
@implementer(interfaces.IGroup)
|
||||
class TwistedWordsGroup(basesupport.AbstractGroup):
|
||||
def __init__(self, name, wordsClient):
|
||||
basesupport.AbstractGroup.__init__(self, name, wordsClient)
|
||||
self.joined = 0
|
||||
|
||||
def sendGroupMessage(self, text, metadata=None):
|
||||
"""Return a deferred.
|
||||
"""
|
||||
#for backwards compatibility with older twisted.words servers.
|
||||
if metadata:
|
||||
d=self.account.client.perspective.callRemote(
|
||||
'groupMessage', self.name, text, metadata)
|
||||
d.addErrback(self.metadataFailed, "* "+text)
|
||||
return d
|
||||
else:
|
||||
return self.account.client.perspective.callRemote('groupMessage',
|
||||
self.name, text)
|
||||
|
||||
def setTopic(self, text):
|
||||
self.account.client.perspective.callRemote(
|
||||
'setGroupMetadata',
|
||||
{'topic': text, 'topic_author': self.client.name},
|
||||
self.name)
|
||||
|
||||
def metadataFailed(self, result, text):
|
||||
print("result:",result,"text:",text)
|
||||
return self.account.client.perspective.callRemote('groupMessage',
|
||||
self.name, text)
|
||||
|
||||
def joining(self):
|
||||
self.joined = 1
|
||||
|
||||
def leaving(self):
|
||||
self.joined = 0
|
||||
|
||||
def leave(self):
|
||||
return self.account.client.perspective.callRemote('leaveGroup',
|
||||
self.name)
|
||||
|
||||
|
||||
|
||||
class TwistedWordsClient(pb.Referenceable, basesupport.AbstractClientMixin):
|
||||
"""In some cases, this acts as an Account, since it a source of text
|
||||
messages (multiple Words instances may be on a single PB connection)
|
||||
"""
|
||||
def __init__(self, acct, serviceName, perspectiveName, chatui,
|
||||
_logonDeferred=None):
|
||||
self.accountName = "%s (%s:%s)" % (acct.accountName, serviceName, perspectiveName)
|
||||
self.name = perspectiveName
|
||||
print("HELLO I AM A PB SERVICE", serviceName, perspectiveName)
|
||||
self.chat = chatui
|
||||
self.account = acct
|
||||
self._logonDeferred = _logonDeferred
|
||||
|
||||
def getPerson(self, name):
|
||||
return self.chat.getPerson(name, self)
|
||||
|
||||
def getGroup(self, name):
|
||||
return self.chat.getGroup(name, self)
|
||||
|
||||
def getGroupConversation(self, name):
|
||||
return self.chat.getGroupConversation(self.getGroup(name))
|
||||
|
||||
def addContact(self, name):
|
||||
self.perspective.callRemote('addContact', name)
|
||||
|
||||
def remote_receiveGroupMembers(self, names, group):
|
||||
print('received group members:', names, group)
|
||||
self.getGroupConversation(group).setGroupMembers(names)
|
||||
|
||||
def remote_receiveGroupMessage(self, sender, group, message, metadata=None):
|
||||
print('received a group message', sender, group, message, metadata)
|
||||
self.getGroupConversation(group).showGroupMessage(sender, message, metadata)
|
||||
|
||||
def remote_memberJoined(self, member, group):
|
||||
print('member joined', member, group)
|
||||
self.getGroupConversation(group).memberJoined(member)
|
||||
|
||||
def remote_memberLeft(self, member, group):
|
||||
print('member left')
|
||||
self.getGroupConversation(group).memberLeft(member)
|
||||
|
||||
def remote_notifyStatusChanged(self, name, status):
|
||||
self.chat.getPerson(name, self).setStatus(status)
|
||||
|
||||
def remote_receiveDirectMessage(self, name, message, metadata=None):
|
||||
self.chat.getConversation(self.chat.getPerson(name, self)).showMessage(message, metadata)
|
||||
|
||||
def remote_receiveContactList(self, clist):
|
||||
for name, status in clist:
|
||||
self.chat.getPerson(name, self).setStatus(status)
|
||||
|
||||
def remote_setGroupMetadata(self, dict_, groupName):
|
||||
if "topic" in dict_:
|
||||
self.getGroupConversation(groupName).setTopic(dict_["topic"], dict_.get("topic_author", None))
|
||||
|
||||
def joinGroup(self, name):
|
||||
self.getGroup(name).joining()
|
||||
return self.perspective.callRemote('joinGroup', name).addCallback(self._cbGroupJoined, name)
|
||||
|
||||
def leaveGroup(self, name):
|
||||
self.getGroup(name).leaving()
|
||||
return self.perspective.callRemote('leaveGroup', name).addCallback(self._cbGroupLeft, name)
|
||||
|
||||
def _cbGroupJoined(self, result, name):
|
||||
groupConv = self.chat.getGroupConversation(self.getGroup(name))
|
||||
groupConv.showGroupMessage("sys", "you joined")
|
||||
self.perspective.callRemote('getGroupMembers', name)
|
||||
|
||||
def _cbGroupLeft(self, result, name):
|
||||
print('left',name)
|
||||
groupConv = self.chat.getGroupConversation(self.getGroup(name), 1)
|
||||
groupConv.showGroupMessage("sys", "you left")
|
||||
|
||||
def connected(self, perspective):
|
||||
print('Connected Words Client!', perspective)
|
||||
if self._logonDeferred is not None:
|
||||
self._logonDeferred.callback(self)
|
||||
self.perspective = perspective
|
||||
self.chat.getContactsList()
|
||||
|
||||
|
||||
pbFrontEnds = {
|
||||
"twisted.words": TwistedWordsClient,
|
||||
"twisted.reality": None
|
||||
}
|
||||
|
||||
|
||||
@implementer(interfaces.IAccount)
|
||||
class PBAccount(basesupport.AbstractAccount):
|
||||
gatewayType = "PB"
|
||||
_groupFactory = TwistedWordsGroup
|
||||
_personFactory = TwistedWordsPerson
|
||||
|
||||
def __init__(self, accountName, autoLogin, username, password, host, port,
|
||||
services=None):
|
||||
"""
|
||||
@param username: The name of your PB Identity.
|
||||
@type username: string
|
||||
"""
|
||||
basesupport.AbstractAccount.__init__(self, accountName, autoLogin,
|
||||
username, password, host, port)
|
||||
self.services = []
|
||||
if not services:
|
||||
services = [('twisted.words', 'twisted.words', username)]
|
||||
for serviceType, serviceName, perspectiveName in services:
|
||||
self.services.append([pbFrontEnds[serviceType], serviceName,
|
||||
perspectiveName])
|
||||
|
||||
def logOn(self, chatui):
|
||||
"""
|
||||
@returns: this breaks with L{interfaces.IAccount}
|
||||
@returntype: DeferredList of L{interfaces.IClient}s
|
||||
"""
|
||||
# Overriding basesupport's implementation on account of the
|
||||
# fact that _startLogOn tends to return a deferredList rather
|
||||
# than a simple Deferred, and we need to do registerAccountClient.
|
||||
if (not self._isConnecting) and (not self._isOnline):
|
||||
self._isConnecting = 1
|
||||
d = self._startLogOn(chatui)
|
||||
d.addErrback(self._loginFailed)
|
||||
def registerMany(results):
|
||||
for success, result in results:
|
||||
if success:
|
||||
chatui.registerAccountClient(result)
|
||||
self._cb_logOn(result)
|
||||
else:
|
||||
log.err(result)
|
||||
d.addCallback(registerMany)
|
||||
return d
|
||||
else:
|
||||
raise error.ConnectionError("Connection in progress")
|
||||
|
||||
|
||||
def _startLogOn(self, chatui):
|
||||
print('Connecting...', end=' ')
|
||||
d = pb.getObjectAt(self.host, self.port)
|
||||
d.addCallbacks(self._cbConnected, self._ebConnected,
|
||||
callbackArgs=(chatui,))
|
||||
return d
|
||||
|
||||
def _cbConnected(self, root, chatui):
|
||||
print('Connected!')
|
||||
print('Identifying...', end=' ')
|
||||
d = pb.authIdentity(root, self.username, self.password)
|
||||
d.addCallbacks(self._cbIdent, self._ebConnected,
|
||||
callbackArgs=(chatui,))
|
||||
return d
|
||||
|
||||
def _cbIdent(self, ident, chatui):
|
||||
if not ident:
|
||||
print('falsely identified.')
|
||||
return self._ebConnected(Failure(Exception("username or password incorrect")))
|
||||
print('Identified!')
|
||||
dl = []
|
||||
for handlerClass, sname, pname in self.services:
|
||||
d = defer.Deferred()
|
||||
dl.append(d)
|
||||
handler = handlerClass(self, sname, pname, chatui, d)
|
||||
ident.callRemote('attach', sname, pname, handler).addCallback(handler.connected)
|
||||
return defer.DeferredList(dl)
|
||||
|
||||
def _ebConnected(self, error):
|
||||
print('Not connected.')
|
||||
return error
|
||||
|
||||
@@ -0,0 +1,267 @@
|
||||
# -*- test-case-name: twisted.words.test -*-
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
from zope.interface import Interface, Attribute
|
||||
|
||||
|
||||
class IProtocolPlugin(Interface):
|
||||
"""Interface for plugins providing an interface to a Words service
|
||||
"""
|
||||
|
||||
name = Attribute("A single word describing what kind of interface this is (eg, irc or web)")
|
||||
|
||||
def getFactory(realm, portal):
|
||||
"""Retrieve a C{twisted.internet.interfaces.IServerFactory} provider
|
||||
|
||||
@param realm: An object providing C{twisted.cred.portal.IRealm} and
|
||||
L{IChatService}, with which service information should be looked up.
|
||||
|
||||
@param portal: An object providing C{twisted.cred.portal.IPortal},
|
||||
through which logins should be performed.
|
||||
"""
|
||||
|
||||
|
||||
class IGroup(Interface):
|
||||
name = Attribute("A short string, unique among groups.")
|
||||
|
||||
def add(user):
|
||||
"""Include the given user in this group.
|
||||
|
||||
@type user: L{IUser}
|
||||
"""
|
||||
|
||||
def remove(user, reason=None):
|
||||
"""Remove the given user from this group.
|
||||
|
||||
@type user: L{IUser}
|
||||
@type reason: C{unicode}
|
||||
"""
|
||||
|
||||
def size():
|
||||
"""Return the number of participants in this group.
|
||||
|
||||
@rtype: L{twisted.internet.defer.Deferred}
|
||||
@return: A Deferred which fires with an C{int} representing the
|
||||
number of participants in this group.
|
||||
"""
|
||||
|
||||
def receive(sender, recipient, message):
|
||||
"""
|
||||
Broadcast the given message from the given sender to other
|
||||
users in group.
|
||||
|
||||
The message is not re-transmitted to the sender.
|
||||
|
||||
@param sender: L{IUser}
|
||||
|
||||
@type recipient: L{IGroup}
|
||||
@param recipient: This is probably a wart. Maybe it will be removed
|
||||
in the future. For now, it should be the group object the message
|
||||
is being delivered to.
|
||||
|
||||
@param message: C{dict}
|
||||
|
||||
@rtype: L{twisted.internet.defer.Deferred}
|
||||
@return: A Deferred which fires with None when delivery has been
|
||||
attempted for all users.
|
||||
"""
|
||||
|
||||
def setMetadata(meta):
|
||||
"""Change the metadata associated with this group.
|
||||
|
||||
@type meta: C{dict}
|
||||
"""
|
||||
|
||||
def iterusers():
|
||||
"""Return an iterator of all users in this group.
|
||||
"""
|
||||
|
||||
|
||||
class IChatClient(Interface):
|
||||
"""Interface through which IChatService interacts with clients.
|
||||
"""
|
||||
|
||||
name = Attribute("A short string, unique among users. This will be set by the L{IChatService} at login time.")
|
||||
|
||||
def receive(sender, recipient, message):
|
||||
"""
|
||||
Callback notifying this user of the given message sent by the
|
||||
given user.
|
||||
|
||||
This will be invoked whenever another user sends a message to a
|
||||
group this user is participating in, or whenever another user sends
|
||||
a message directly to this user. In the former case, C{recipient}
|
||||
will be the group to which the message was sent; in the latter, it
|
||||
will be the same object as the user who is receiving the message.
|
||||
|
||||
@type sender: L{IUser}
|
||||
@type recipient: L{IUser} or L{IGroup}
|
||||
@type message: C{dict}
|
||||
|
||||
@rtype: L{twisted.internet.defer.Deferred}
|
||||
@return: A Deferred which fires when the message has been delivered,
|
||||
or which fails in some way. If the Deferred fails and the message
|
||||
was directed at a group, this user will be removed from that group.
|
||||
"""
|
||||
|
||||
def groupMetaUpdate(group, meta):
|
||||
"""
|
||||
Callback notifying this user that the metadata for the given
|
||||
group has changed.
|
||||
|
||||
@type group: L{IGroup}
|
||||
@type meta: C{dict}
|
||||
|
||||
@rtype: L{twisted.internet.defer.Deferred}
|
||||
"""
|
||||
|
||||
def userJoined(group, user):
|
||||
"""
|
||||
Callback notifying this user that the given user has joined
|
||||
the given group.
|
||||
|
||||
@type group: L{IGroup}
|
||||
@type user: L{IUser}
|
||||
|
||||
@rtype: L{twisted.internet.defer.Deferred}
|
||||
"""
|
||||
|
||||
def userLeft(group, user, reason=None):
|
||||
"""
|
||||
Callback notifying this user that the given user has left the
|
||||
given group for the given reason.
|
||||
|
||||
@type group: L{IGroup}
|
||||
@type user: L{IUser}
|
||||
@type reason: C{unicode}
|
||||
|
||||
@rtype: L{twisted.internet.defer.Deferred}
|
||||
"""
|
||||
|
||||
|
||||
class IUser(Interface):
|
||||
"""Interface through which clients interact with IChatService.
|
||||
"""
|
||||
|
||||
realm = Attribute("A reference to the Realm to which this user belongs. Set if and only if the user is logged in.")
|
||||
mind = Attribute("A reference to the mind which logged in to this user. Set if and only if the user is logged in.")
|
||||
name = Attribute("A short string, unique among users.")
|
||||
|
||||
lastMessage = Attribute("A POSIX timestamp indicating the time of the last message received from this user.")
|
||||
signOn = Attribute("A POSIX timestamp indicating this user's most recent sign on time.")
|
||||
|
||||
def loggedIn(realm, mind):
|
||||
"""Invoked by the associated L{IChatService} when login occurs.
|
||||
|
||||
@param realm: The L{IChatService} through which login is occurring.
|
||||
@param mind: The mind object used for cred login.
|
||||
"""
|
||||
|
||||
def send(recipient, message):
|
||||
"""Send the given message to the given user or group.
|
||||
|
||||
@type recipient: Either L{IUser} or L{IGroup}
|
||||
@type message: C{dict}
|
||||
"""
|
||||
|
||||
def join(group):
|
||||
"""Attempt to join the given group.
|
||||
|
||||
@type group: L{IGroup}
|
||||
@rtype: L{twisted.internet.defer.Deferred}
|
||||
"""
|
||||
|
||||
def leave(group):
|
||||
"""Discontinue participation in the given group.
|
||||
|
||||
@type group: L{IGroup}
|
||||
@rtype: L{twisted.internet.defer.Deferred}
|
||||
"""
|
||||
|
||||
def itergroups():
|
||||
"""
|
||||
Return an iterator of all groups of which this user is a
|
||||
member.
|
||||
"""
|
||||
|
||||
|
||||
class IChatService(Interface):
|
||||
name = Attribute("A short string identifying this chat service (eg, a hostname)")
|
||||
|
||||
createGroupOnRequest = Attribute(
|
||||
"A boolean indicating whether L{getGroup} should implicitly "
|
||||
"create groups which are requested but which do not yet exist.")
|
||||
|
||||
createUserOnRequest = Attribute(
|
||||
"A boolean indicating whether L{getUser} should implicitly "
|
||||
"create users which are requested but which do not yet exist.")
|
||||
|
||||
def itergroups():
|
||||
"""Return all groups available on this service.
|
||||
|
||||
@rtype: C{twisted.internet.defer.Deferred}
|
||||
@return: A Deferred which fires with a list of C{IGroup} providers.
|
||||
"""
|
||||
|
||||
def getGroup(name):
|
||||
"""Retrieve the group by the given name.
|
||||
|
||||
@type name: C{str}
|
||||
|
||||
@rtype: L{twisted.internet.defer.Deferred}
|
||||
@return: A Deferred which fires with the group with the given
|
||||
name if one exists (or if one is created due to the setting of
|
||||
L{IChatService.createGroupOnRequest}, or which fails with
|
||||
L{twisted.words.ewords.NoSuchGroup} if no such group exists.
|
||||
"""
|
||||
|
||||
def createGroup(name):
|
||||
"""Create a new group with the given name.
|
||||
|
||||
@type name: C{str}
|
||||
|
||||
@rtype: L{twisted.internet.defer.Deferred}
|
||||
@return: A Deferred which fires with the created group, or
|
||||
with fails with L{twisted.words.ewords.DuplicateGroup} if a
|
||||
group by that name exists already.
|
||||
"""
|
||||
|
||||
def lookupGroup(name):
|
||||
"""Retrieve a group by name.
|
||||
|
||||
Unlike C{getGroup}, this will never implicitly create a group.
|
||||
|
||||
@type name: C{str}
|
||||
|
||||
@rtype: L{twisted.internet.defer.Deferred}
|
||||
@return: A Deferred which fires with the group by the given
|
||||
name, or which fails with L{twisted.words.ewords.NoSuchGroup}.
|
||||
"""
|
||||
|
||||
def getUser(name):
|
||||
"""Retrieve the user by the given name.
|
||||
|
||||
@type name: C{str}
|
||||
|
||||
@rtype: L{twisted.internet.defer.Deferred}
|
||||
@return: A Deferred which fires with the user with the given
|
||||
name if one exists (or if one is created due to the setting of
|
||||
L{IChatService.createUserOnRequest}, or which fails with
|
||||
L{twisted.words.ewords.NoSuchUser} if no such user exists.
|
||||
"""
|
||||
|
||||
def createUser(name):
|
||||
"""Create a new user with the given name.
|
||||
|
||||
@type name: C{str}
|
||||
|
||||
@rtype: L{twisted.internet.defer.Deferred}
|
||||
@return: A Deferred which fires with the created user, or
|
||||
with fails with L{twisted.words.ewords.DuplicateUser} if a
|
||||
user by that name exists already.
|
||||
"""
|
||||
|
||||
__all__ = [
|
||||
'IGroup', 'IChatClient', 'IUser', 'IChatService',
|
||||
]
|
||||
@@ -0,0 +1,233 @@
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
XMPP-specific SASL profile.
|
||||
"""
|
||||
|
||||
from __future__ import absolute_import, division
|
||||
|
||||
from base64 import b64decode, b64encode
|
||||
import re
|
||||
from twisted.internet import defer
|
||||
from twisted.python.compat import unicode
|
||||
from twisted.words.protocols.jabber import sasl_mechanisms, xmlstream
|
||||
from twisted.words.xish import domish
|
||||
|
||||
NS_XMPP_SASL = 'urn:ietf:params:xml:ns:xmpp-sasl'
|
||||
|
||||
def get_mechanisms(xs):
|
||||
"""
|
||||
Parse the SASL feature to extract the available mechanism names.
|
||||
"""
|
||||
mechanisms = []
|
||||
for element in xs.features[(NS_XMPP_SASL, 'mechanisms')].elements():
|
||||
if element.name == 'mechanism':
|
||||
mechanisms.append(unicode(element))
|
||||
|
||||
return mechanisms
|
||||
|
||||
|
||||
class SASLError(Exception):
|
||||
"""
|
||||
SASL base exception.
|
||||
"""
|
||||
|
||||
|
||||
class SASLNoAcceptableMechanism(SASLError):
|
||||
"""
|
||||
The server did not present an acceptable SASL mechanism.
|
||||
"""
|
||||
|
||||
|
||||
class SASLAuthError(SASLError):
|
||||
"""
|
||||
SASL Authentication failed.
|
||||
"""
|
||||
def __init__(self, condition=None):
|
||||
self.condition = condition
|
||||
|
||||
|
||||
def __str__(self):
|
||||
return "SASLAuthError with condition %r" % self.condition
|
||||
|
||||
|
||||
class SASLIncorrectEncodingError(SASLError):
|
||||
"""
|
||||
SASL base64 encoding was incorrect.
|
||||
|
||||
RFC 3920 specifies that any characters not in the base64 alphabet
|
||||
and padding characters present elsewhere than at the end of the string
|
||||
MUST be rejected. See also L{fromBase64}.
|
||||
|
||||
This exception is raised whenever the encoded string does not adhere
|
||||
to these additional restrictions or when the decoding itself fails.
|
||||
|
||||
The recommended behaviour for so-called receiving entities (like servers in
|
||||
client-to-server connections, see RFC 3920 for terminology) is to fail the
|
||||
SASL negotiation with a C{'incorrect-encoding'} condition. For initiating
|
||||
entities, one should assume the receiving entity to be either buggy or
|
||||
malevolent. The stream should be terminated and reconnecting is not
|
||||
advised.
|
||||
"""
|
||||
|
||||
base64Pattern = re.compile("^[0-9A-Za-z+/]*[0-9A-Za-z+/=]{,2}$")
|
||||
|
||||
def fromBase64(s):
|
||||
"""
|
||||
Decode base64 encoded string.
|
||||
|
||||
This helper performs regular decoding of a base64 encoded string, but also
|
||||
rejects any characters that are not in the base64 alphabet and padding
|
||||
occurring elsewhere from the last or last two characters, as specified in
|
||||
section 14.9 of RFC 3920. This safeguards against various attack vectors
|
||||
among which the creation of a covert channel that "leaks" information.
|
||||
"""
|
||||
|
||||
if base64Pattern.match(s) is None:
|
||||
raise SASLIncorrectEncodingError()
|
||||
|
||||
try:
|
||||
return b64decode(s)
|
||||
except Exception as e:
|
||||
raise SASLIncorrectEncodingError(str(e))
|
||||
|
||||
|
||||
|
||||
class SASLInitiatingInitializer(xmlstream.BaseFeatureInitiatingInitializer):
|
||||
"""
|
||||
Stream initializer that performs SASL authentication.
|
||||
|
||||
The supported mechanisms by this initializer are C{DIGEST-MD5}, C{PLAIN}
|
||||
and C{ANONYMOUS}. The C{ANONYMOUS} SASL mechanism is used when the JID, set
|
||||
on the authenticator, does not have a localpart (username), requesting an
|
||||
anonymous session where the username is generated by the server.
|
||||
Otherwise, C{DIGEST-MD5} and C{PLAIN} are attempted, in that order.
|
||||
"""
|
||||
|
||||
feature = (NS_XMPP_SASL, 'mechanisms')
|
||||
_deferred = None
|
||||
|
||||
def setMechanism(self):
|
||||
"""
|
||||
Select and setup authentication mechanism.
|
||||
|
||||
Uses the authenticator's C{jid} and C{password} attribute for the
|
||||
authentication credentials. If no supported SASL mechanisms are
|
||||
advertized by the receiving party, a failing deferred is returned with
|
||||
a L{SASLNoAcceptableMechanism} exception.
|
||||
"""
|
||||
|
||||
jid = self.xmlstream.authenticator.jid
|
||||
password = self.xmlstream.authenticator.password
|
||||
|
||||
mechanisms = get_mechanisms(self.xmlstream)
|
||||
if jid.user is not None:
|
||||
if 'DIGEST-MD5' in mechanisms:
|
||||
self.mechanism = sasl_mechanisms.DigestMD5('xmpp', jid.host, None,
|
||||
jid.user, password)
|
||||
elif 'PLAIN' in mechanisms:
|
||||
self.mechanism = sasl_mechanisms.Plain(None, jid.user, password)
|
||||
else:
|
||||
raise SASLNoAcceptableMechanism()
|
||||
else:
|
||||
if 'ANONYMOUS' in mechanisms:
|
||||
self.mechanism = sasl_mechanisms.Anonymous()
|
||||
else:
|
||||
raise SASLNoAcceptableMechanism()
|
||||
|
||||
|
||||
def start(self):
|
||||
"""
|
||||
Start SASL authentication exchange.
|
||||
"""
|
||||
|
||||
self.setMechanism()
|
||||
self._deferred = defer.Deferred()
|
||||
self.xmlstream.addObserver('/challenge', self.onChallenge)
|
||||
self.xmlstream.addOnetimeObserver('/success', self.onSuccess)
|
||||
self.xmlstream.addOnetimeObserver('/failure', self.onFailure)
|
||||
self.sendAuth(self.mechanism.getInitialResponse())
|
||||
return self._deferred
|
||||
|
||||
|
||||
def sendAuth(self, data=None):
|
||||
"""
|
||||
Initiate authentication protocol exchange.
|
||||
|
||||
If an initial client response is given in C{data}, it will be
|
||||
sent along.
|
||||
|
||||
@param data: initial client response.
|
||||
@type data: C{str} or L{None}.
|
||||
"""
|
||||
|
||||
auth = domish.Element((NS_XMPP_SASL, 'auth'))
|
||||
auth['mechanism'] = self.mechanism.name
|
||||
if data is not None:
|
||||
auth.addContent(b64encode(data).decode('ascii') or u'=')
|
||||
self.xmlstream.send(auth)
|
||||
|
||||
|
||||
def sendResponse(self, data=b''):
|
||||
"""
|
||||
Send response to a challenge.
|
||||
|
||||
@param data: client response.
|
||||
@type data: L{bytes}.
|
||||
"""
|
||||
|
||||
response = domish.Element((NS_XMPP_SASL, 'response'))
|
||||
if data:
|
||||
response.addContent(b64encode(data).decode('ascii'))
|
||||
self.xmlstream.send(response)
|
||||
|
||||
|
||||
def onChallenge(self, element):
|
||||
"""
|
||||
Parse challenge and send response from the mechanism.
|
||||
|
||||
@param element: the challenge protocol element.
|
||||
@type element: L{domish.Element}.
|
||||
"""
|
||||
|
||||
try:
|
||||
challenge = fromBase64(unicode(element))
|
||||
except SASLIncorrectEncodingError:
|
||||
self._deferred.errback()
|
||||
else:
|
||||
self.sendResponse(self.mechanism.getResponse(challenge))
|
||||
|
||||
|
||||
def onSuccess(self, success):
|
||||
"""
|
||||
Clean up observers, reset the XML stream and send a new header.
|
||||
|
||||
@param success: the success protocol element. For now unused, but
|
||||
could hold additional data.
|
||||
@type success: L{domish.Element}
|
||||
"""
|
||||
|
||||
self.xmlstream.removeObserver('/challenge', self.onChallenge)
|
||||
self.xmlstream.removeObserver('/failure', self.onFailure)
|
||||
self.xmlstream.reset()
|
||||
self.xmlstream.sendHeader()
|
||||
self._deferred.callback(xmlstream.Reset)
|
||||
|
||||
|
||||
def onFailure(self, failure):
|
||||
"""
|
||||
Clean up observers, parse the failure and errback the deferred.
|
||||
|
||||
@param failure: the failure protocol element. Holds details on
|
||||
the error condition.
|
||||
@type failure: L{domish.Element}
|
||||
"""
|
||||
|
||||
self.xmlstream.removeObserver('/challenge', self.onChallenge)
|
||||
self.xmlstream.removeObserver('/success', self.onSuccess)
|
||||
try:
|
||||
condition = failure.firstChildElement().name
|
||||
except AttributeError:
|
||||
condition = None
|
||||
self._deferred.errback(SASLAuthError(condition))
|
||||
@@ -0,0 +1,293 @@
|
||||
# -*- test-case-name: twisted.words.test.test_jabbersaslmechanisms -*-
|
||||
#
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
Protocol agnostic implementations of SASL authentication mechanisms.
|
||||
"""
|
||||
|
||||
from __future__ import absolute_import, division
|
||||
|
||||
import binascii, random, time, os
|
||||
from hashlib import md5
|
||||
|
||||
from zope.interface import Interface, Attribute, implementer
|
||||
|
||||
from twisted.python.compat import iteritems, networkString
|
||||
|
||||
|
||||
class ISASLMechanism(Interface):
|
||||
name = Attribute("""Common name for the SASL Mechanism.""")
|
||||
|
||||
def getInitialResponse():
|
||||
"""
|
||||
Get the initial client response, if defined for this mechanism.
|
||||
|
||||
@return: initial client response string.
|
||||
@rtype: C{str}.
|
||||
"""
|
||||
|
||||
|
||||
def getResponse(challenge):
|
||||
"""
|
||||
Get the response to a server challenge.
|
||||
|
||||
@param challenge: server challenge.
|
||||
@type challenge: C{str}.
|
||||
@return: client response.
|
||||
@rtype: C{str}.
|
||||
"""
|
||||
|
||||
|
||||
|
||||
@implementer(ISASLMechanism)
|
||||
class Anonymous(object):
|
||||
"""
|
||||
Implements the ANONYMOUS SASL authentication mechanism.
|
||||
|
||||
This mechanism is defined in RFC 2245.
|
||||
"""
|
||||
name = 'ANONYMOUS'
|
||||
|
||||
def getInitialResponse(self):
|
||||
return None
|
||||
|
||||
|
||||
|
||||
@implementer(ISASLMechanism)
|
||||
class Plain(object):
|
||||
"""
|
||||
Implements the PLAIN SASL authentication mechanism.
|
||||
|
||||
The PLAIN SASL authentication mechanism is defined in RFC 2595.
|
||||
"""
|
||||
name = 'PLAIN'
|
||||
|
||||
def __init__(self, authzid, authcid, password):
|
||||
"""
|
||||
@param authzid: The authorization identity.
|
||||
@type authzid: L{unicode}
|
||||
|
||||
@param authcid: The authentication identity.
|
||||
@type authcid: L{unicode}
|
||||
|
||||
@param password: The plain-text password.
|
||||
@type password: L{unicode}
|
||||
"""
|
||||
|
||||
self.authzid = authzid or u''
|
||||
self.authcid = authcid or u''
|
||||
self.password = password or u''
|
||||
|
||||
|
||||
def getInitialResponse(self):
|
||||
return (self.authzid.encode('utf-8') + b"\x00" +
|
||||
self.authcid.encode('utf-8') + b"\x00" +
|
||||
self.password.encode('utf-8'))
|
||||
|
||||
|
||||
|
||||
@implementer(ISASLMechanism)
|
||||
class DigestMD5(object):
|
||||
"""
|
||||
Implements the DIGEST-MD5 SASL authentication mechanism.
|
||||
|
||||
The DIGEST-MD5 SASL authentication mechanism is defined in RFC 2831.
|
||||
"""
|
||||
name = 'DIGEST-MD5'
|
||||
|
||||
def __init__(self, serv_type, host, serv_name, username, password):
|
||||
"""
|
||||
@param serv_type: An indication of what kind of server authentication
|
||||
is being attempted against. For example, C{u"xmpp"}.
|
||||
@type serv_type: C{unicode}
|
||||
|
||||
@param host: The authentication hostname. Also known as the realm.
|
||||
This is used as a scope to help select the right credentials.
|
||||
@type host: C{unicode}
|
||||
|
||||
@param serv_name: An additional identifier for the server.
|
||||
@type serv_name: C{unicode}
|
||||
|
||||
@param username: The authentication username to use to respond to a
|
||||
challenge.
|
||||
@type username: C{unicode}
|
||||
|
||||
@param username: The authentication password to use to respond to a
|
||||
challenge.
|
||||
@type password: C{unicode}
|
||||
"""
|
||||
self.username = username
|
||||
self.password = password
|
||||
self.defaultRealm = host
|
||||
|
||||
self.digest_uri = u'%s/%s' % (serv_type, host)
|
||||
if serv_name is not None:
|
||||
self.digest_uri += u'/%s' % (serv_name,)
|
||||
|
||||
|
||||
def getInitialResponse(self):
|
||||
return None
|
||||
|
||||
|
||||
def getResponse(self, challenge):
|
||||
directives = self._parse(challenge)
|
||||
|
||||
# Compat for implementations that do not send this along with
|
||||
# a successful authentication.
|
||||
if b'rspauth' in directives:
|
||||
return b''
|
||||
|
||||
charset = directives[b'charset'].decode('ascii')
|
||||
|
||||
try:
|
||||
realm = directives[b'realm']
|
||||
except KeyError:
|
||||
realm = self.defaultRealm.encode(charset)
|
||||
|
||||
return self._genResponse(charset,
|
||||
realm,
|
||||
directives[b'nonce'])
|
||||
|
||||
|
||||
def _parse(self, challenge):
|
||||
"""
|
||||
Parses the server challenge.
|
||||
|
||||
Splits the challenge into a dictionary of directives with values.
|
||||
|
||||
@return: challenge directives and their values.
|
||||
@rtype: C{dict} of C{str} to C{str}.
|
||||
"""
|
||||
s = challenge
|
||||
paramDict = {}
|
||||
cur = 0
|
||||
remainingParams = True
|
||||
while remainingParams:
|
||||
# Parse a param. We can't just split on commas, because there can
|
||||
# be some commas inside (quoted) param values, e.g.:
|
||||
# qop="auth,auth-int"
|
||||
|
||||
middle = s.index(b"=", cur)
|
||||
name = s[cur:middle].lstrip()
|
||||
middle += 1
|
||||
if s[middle:middle+1] == b'"':
|
||||
middle += 1
|
||||
end = s.index(b'"', middle)
|
||||
value = s[middle:end]
|
||||
cur = s.find(b',', end) + 1
|
||||
if cur == 0:
|
||||
remainingParams = False
|
||||
else:
|
||||
end = s.find(b',', middle)
|
||||
if end == -1:
|
||||
value = s[middle:].rstrip()
|
||||
remainingParams = False
|
||||
else:
|
||||
value = s[middle:end].rstrip()
|
||||
cur = end + 1
|
||||
paramDict[name] = value
|
||||
|
||||
for param in (b'qop', b'cipher'):
|
||||
if param in paramDict:
|
||||
paramDict[param] = paramDict[param].split(b',')
|
||||
|
||||
return paramDict
|
||||
|
||||
def _unparse(self, directives):
|
||||
"""
|
||||
Create message string from directives.
|
||||
|
||||
@param directives: dictionary of directives (names to their values).
|
||||
For certain directives, extra quotes are added, as
|
||||
needed.
|
||||
@type directives: C{dict} of C{str} to C{str}
|
||||
@return: message string.
|
||||
@rtype: C{str}.
|
||||
"""
|
||||
|
||||
directive_list = []
|
||||
for name, value in iteritems(directives):
|
||||
if name in (b'username', b'realm', b'cnonce',
|
||||
b'nonce', b'digest-uri', b'authzid', b'cipher'):
|
||||
directive = name + b'=' + value
|
||||
else:
|
||||
directive = name + b'=' + value
|
||||
|
||||
directive_list.append(directive)
|
||||
|
||||
return b','.join(directive_list)
|
||||
|
||||
|
||||
def _calculateResponse(self, cnonce, nc, nonce,
|
||||
username, password, realm, uri):
|
||||
"""
|
||||
Calculates response with given encoded parameters.
|
||||
|
||||
@return: The I{response} field of a response to a Digest-MD5 challenge
|
||||
of the given parameters.
|
||||
@rtype: L{bytes}
|
||||
"""
|
||||
def H(s):
|
||||
return md5(s).digest()
|
||||
|
||||
def HEX(n):
|
||||
return binascii.b2a_hex(n)
|
||||
|
||||
def KD(k, s):
|
||||
return H(k + b':' + s)
|
||||
|
||||
a1 = (H(username + b":" + realm + b":" + password) + b":" +
|
||||
nonce + b":" +
|
||||
cnonce)
|
||||
a2 = b"AUTHENTICATE:" + uri
|
||||
|
||||
response = HEX(KD(HEX(H(a1)),
|
||||
nonce + b":" + nc + b":" + cnonce + b":" +
|
||||
b"auth" + b":" + HEX(H(a2))))
|
||||
return response
|
||||
|
||||
|
||||
def _genResponse(self, charset, realm, nonce):
|
||||
"""
|
||||
Generate response-value.
|
||||
|
||||
Creates a response to a challenge according to section 2.1.2.1 of
|
||||
RFC 2831 using the C{charset}, C{realm} and C{nonce} directives
|
||||
from the challenge.
|
||||
"""
|
||||
try:
|
||||
username = self.username.encode(charset)
|
||||
password = self.password.encode(charset)
|
||||
digest_uri = self.digest_uri.encode(charset)
|
||||
except UnicodeError:
|
||||
# TODO - add error checking
|
||||
raise
|
||||
|
||||
nc = networkString('%08x' % (1,)) # TODO: support subsequent auth.
|
||||
cnonce = self._gen_nonce()
|
||||
qop = b'auth'
|
||||
|
||||
# TODO - add support for authzid
|
||||
response = self._calculateResponse(cnonce, nc, nonce,
|
||||
username, password, realm,
|
||||
digest_uri)
|
||||
|
||||
directives = {b'username': username,
|
||||
b'realm' : realm,
|
||||
b'nonce' : nonce,
|
||||
b'cnonce' : cnonce,
|
||||
b'nc' : nc,
|
||||
b'qop' : qop,
|
||||
b'digest-uri': digest_uri,
|
||||
b'response': response,
|
||||
b'charset': charset.encode('ascii')}
|
||||
|
||||
return self._unparse(directives)
|
||||
|
||||
|
||||
def _gen_nonce(self):
|
||||
nonceString = "%f:%f:%d" % (random.random(), time.time(), os.getpid())
|
||||
nonceBytes = networkString(nonceString)
|
||||
return md5(nonceBytes).hexdigest().encode('ascii')
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,244 @@
|
||||
# -*- test-case-name: twisted.words.test.test_jabberxmppstringprep -*-
|
||||
#
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
from encodings import idna
|
||||
from itertools import chain
|
||||
import stringprep
|
||||
|
||||
# We require Unicode version 3.2.
|
||||
from unicodedata import ucd_3_2_0 as unicodedata
|
||||
|
||||
from twisted.python.compat import unichr
|
||||
from twisted.python.deprecate import deprecatedModuleAttribute
|
||||
from incremental import Version
|
||||
|
||||
from zope.interface import Interface, implementer
|
||||
|
||||
|
||||
crippled = False
|
||||
deprecatedModuleAttribute(
|
||||
Version("Twisted", 13, 1, 0),
|
||||
"crippled is always False",
|
||||
__name__,
|
||||
"crippled")
|
||||
|
||||
|
||||
|
||||
class ILookupTable(Interface):
|
||||
"""
|
||||
Interface for character lookup classes.
|
||||
"""
|
||||
|
||||
def lookup(c):
|
||||
"""
|
||||
Return whether character is in this table.
|
||||
"""
|
||||
|
||||
|
||||
|
||||
class IMappingTable(Interface):
|
||||
"""
|
||||
Interface for character mapping classes.
|
||||
"""
|
||||
|
||||
def map(c):
|
||||
"""
|
||||
Return mapping for character.
|
||||
"""
|
||||
|
||||
|
||||
|
||||
@implementer(ILookupTable)
|
||||
class LookupTableFromFunction:
|
||||
|
||||
def __init__(self, in_table_function):
|
||||
self.lookup = in_table_function
|
||||
|
||||
|
||||
|
||||
@implementer(ILookupTable)
|
||||
class LookupTable:
|
||||
|
||||
def __init__(self, table):
|
||||
self._table = table
|
||||
|
||||
def lookup(self, c):
|
||||
return c in self._table
|
||||
|
||||
|
||||
|
||||
@implementer(IMappingTable)
|
||||
class MappingTableFromFunction:
|
||||
|
||||
def __init__(self, map_table_function):
|
||||
self.map = map_table_function
|
||||
|
||||
|
||||
|
||||
@implementer(IMappingTable)
|
||||
class EmptyMappingTable:
|
||||
|
||||
def __init__(self, in_table_function):
|
||||
self._in_table_function = in_table_function
|
||||
|
||||
def map(self, c):
|
||||
if self._in_table_function(c):
|
||||
return None
|
||||
else:
|
||||
return c
|
||||
|
||||
|
||||
|
||||
class Profile:
|
||||
def __init__(self, mappings=[], normalize=True, prohibiteds=[],
|
||||
check_unassigneds=True, check_bidi=True):
|
||||
self.mappings = mappings
|
||||
self.normalize = normalize
|
||||
self.prohibiteds = prohibiteds
|
||||
self.do_check_unassigneds = check_unassigneds
|
||||
self.do_check_bidi = check_bidi
|
||||
|
||||
def prepare(self, string):
|
||||
result = self.map(string)
|
||||
if self.normalize:
|
||||
result = unicodedata.normalize("NFKC", result)
|
||||
self.check_prohibiteds(result)
|
||||
if self.do_check_unassigneds:
|
||||
self.check_unassigneds(result)
|
||||
if self.do_check_bidi:
|
||||
self.check_bidirectionals(result)
|
||||
return result
|
||||
|
||||
def map(self, string):
|
||||
result = []
|
||||
|
||||
for c in string:
|
||||
result_c = c
|
||||
|
||||
for mapping in self.mappings:
|
||||
result_c = mapping.map(c)
|
||||
if result_c != c:
|
||||
break
|
||||
|
||||
if result_c is not None:
|
||||
result.append(result_c)
|
||||
|
||||
return u"".join(result)
|
||||
|
||||
def check_prohibiteds(self, string):
|
||||
for c in string:
|
||||
for table in self.prohibiteds:
|
||||
if table.lookup(c):
|
||||
raise UnicodeError("Invalid character %s" % repr(c))
|
||||
|
||||
def check_unassigneds(self, string):
|
||||
for c in string:
|
||||
if stringprep.in_table_a1(c):
|
||||
raise UnicodeError("Unassigned code point %s" % repr(c))
|
||||
|
||||
def check_bidirectionals(self, string):
|
||||
found_LCat = False
|
||||
found_RandALCat = False
|
||||
|
||||
for c in string:
|
||||
if stringprep.in_table_d1(c):
|
||||
found_RandALCat = True
|
||||
if stringprep.in_table_d2(c):
|
||||
found_LCat = True
|
||||
|
||||
if found_LCat and found_RandALCat:
|
||||
raise UnicodeError("Violation of BIDI Requirement 2")
|
||||
|
||||
if found_RandALCat and not (stringprep.in_table_d1(string[0]) and
|
||||
stringprep.in_table_d1(string[-1])):
|
||||
raise UnicodeError("Violation of BIDI Requirement 3")
|
||||
|
||||
|
||||
class NamePrep:
|
||||
""" Implements preparation of internationalized domain names.
|
||||
|
||||
This class implements preparing internationalized domain names using the
|
||||
rules defined in RFC 3491, section 4 (Conversion operations).
|
||||
|
||||
We do not perform step 4 since we deal with unicode representations of
|
||||
domain names and do not convert from or to ASCII representations using
|
||||
punycode encoding. When such a conversion is needed, the C{idna} standard
|
||||
library provides the C{ToUnicode()} and C{ToASCII()} functions. Note that
|
||||
C{idna} itself assumes UseSTD3ASCIIRules to be false.
|
||||
|
||||
The following steps are performed by C{prepare()}:
|
||||
|
||||
- Split the domain name in labels at the dots (RFC 3490, 3.1)
|
||||
- Apply nameprep proper on each label (RFC 3491)
|
||||
- Enforce the restrictions on ASCII characters in host names by
|
||||
assuming STD3ASCIIRules to be true. (STD 3)
|
||||
- Rejoin the labels using the label separator U+002E (full stop).
|
||||
|
||||
"""
|
||||
|
||||
# Prohibited characters.
|
||||
prohibiteds = [unichr(n) for n in chain(range(0x00, 0x2c + 1),
|
||||
range(0x2e, 0x2f + 1),
|
||||
range(0x3a, 0x40 + 1),
|
||||
range(0x5b, 0x60 + 1),
|
||||
range(0x7b, 0x7f + 1))]
|
||||
|
||||
def prepare(self, string):
|
||||
result = []
|
||||
|
||||
labels = idna.dots.split(string)
|
||||
|
||||
if labels and len(labels[-1]) == 0:
|
||||
trailing_dot = u'.'
|
||||
del labels[-1]
|
||||
else:
|
||||
trailing_dot = u''
|
||||
|
||||
for label in labels:
|
||||
result.append(self.nameprep(label))
|
||||
|
||||
return u".".join(result) + trailing_dot
|
||||
|
||||
def check_prohibiteds(self, string):
|
||||
for c in string:
|
||||
if c in self.prohibiteds:
|
||||
raise UnicodeError("Invalid character %s" % repr(c))
|
||||
|
||||
def nameprep(self, label):
|
||||
label = idna.nameprep(label)
|
||||
self.check_prohibiteds(label)
|
||||
if label[0] == u'-':
|
||||
raise UnicodeError("Invalid leading hyphen-minus")
|
||||
if label[-1] == u'-':
|
||||
raise UnicodeError("Invalid trailing hyphen-minus")
|
||||
return label
|
||||
|
||||
|
||||
C_11 = LookupTableFromFunction(stringprep.in_table_c11)
|
||||
C_12 = LookupTableFromFunction(stringprep.in_table_c12)
|
||||
C_21 = LookupTableFromFunction(stringprep.in_table_c21)
|
||||
C_22 = LookupTableFromFunction(stringprep.in_table_c22)
|
||||
C_3 = LookupTableFromFunction(stringprep.in_table_c3)
|
||||
C_4 = LookupTableFromFunction(stringprep.in_table_c4)
|
||||
C_5 = LookupTableFromFunction(stringprep.in_table_c5)
|
||||
C_6 = LookupTableFromFunction(stringprep.in_table_c6)
|
||||
C_7 = LookupTableFromFunction(stringprep.in_table_c7)
|
||||
C_8 = LookupTableFromFunction(stringprep.in_table_c8)
|
||||
C_9 = LookupTableFromFunction(stringprep.in_table_c9)
|
||||
|
||||
B_1 = EmptyMappingTable(stringprep.in_table_b1)
|
||||
B_2 = MappingTableFromFunction(stringprep.map_table_b2)
|
||||
|
||||
nodeprep = Profile(mappings=[B_1, B_2],
|
||||
prohibiteds=[C_11, C_12, C_21, C_22,
|
||||
C_3, C_4, C_5, C_6, C_7, C_8, C_9,
|
||||
LookupTable([u'"', u'&', u"'", u'/',
|
||||
u':', u'<', u'>', u'@'])])
|
||||
|
||||
resourceprep = Profile(mappings=[B_1,],
|
||||
prohibiteds=[C_12, C_21, C_22,
|
||||
C_3, C_4, C_5, C_6, C_7, C_8, C_9])
|
||||
|
||||
nameprep = NamePrep()
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1 @@
|
||||
"Words tests"
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,440 @@
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
Tests for L{twisted.words.protocols.jabber.component}
|
||||
"""
|
||||
from hashlib import sha1
|
||||
|
||||
from zope.interface.verify import verifyObject
|
||||
|
||||
from twisted.python import failure
|
||||
from twisted.python.compat import unicode
|
||||
from twisted.trial import unittest
|
||||
from twisted.words.protocols.jabber import component, ijabber, xmlstream
|
||||
from twisted.words.protocols.jabber.jid import JID
|
||||
from twisted.words.xish import domish
|
||||
from twisted.words.xish.utility import XmlPipe
|
||||
|
||||
class DummyTransport:
|
||||
def __init__(self, list):
|
||||
self.list = list
|
||||
|
||||
def write(self, bytes):
|
||||
self.list.append(bytes)
|
||||
|
||||
class ComponentInitiatingInitializerTests(unittest.TestCase):
|
||||
def setUp(self):
|
||||
self.output = []
|
||||
|
||||
self.authenticator = xmlstream.Authenticator()
|
||||
self.authenticator.password = u'secret'
|
||||
self.xmlstream = xmlstream.XmlStream(self.authenticator)
|
||||
self.xmlstream.namespace = 'test:component'
|
||||
self.xmlstream.send = self.output.append
|
||||
self.xmlstream.connectionMade()
|
||||
self.xmlstream.dataReceived(
|
||||
"<stream:stream xmlns='test:component' "
|
||||
"xmlns:stream='http://etherx.jabber.org/streams' "
|
||||
"from='example.com' id='12345' version='1.0'>")
|
||||
self.xmlstream.sid = u'12345'
|
||||
self.init = component.ComponentInitiatingInitializer(self.xmlstream)
|
||||
|
||||
def testHandshake(self):
|
||||
"""
|
||||
Test basic operations of component handshake.
|
||||
"""
|
||||
|
||||
d = self.init.initialize()
|
||||
|
||||
# the initializer should have sent the handshake request
|
||||
|
||||
handshake = self.output[-1]
|
||||
self.assertEqual('handshake', handshake.name)
|
||||
self.assertEqual('test:component', handshake.uri)
|
||||
self.assertEqual(sha1(b'12345' + b'secret').hexdigest(),
|
||||
unicode(handshake))
|
||||
|
||||
# successful authentication
|
||||
|
||||
handshake.children = []
|
||||
self.xmlstream.dataReceived(handshake.toXml())
|
||||
|
||||
return d
|
||||
|
||||
class ComponentAuthTests(unittest.TestCase):
|
||||
def authPassed(self, stream):
|
||||
self.authComplete = True
|
||||
|
||||
def testAuth(self):
|
||||
self.authComplete = False
|
||||
outlist = []
|
||||
|
||||
ca = component.ConnectComponentAuthenticator(u"cjid", u"secret")
|
||||
xs = xmlstream.XmlStream(ca)
|
||||
xs.transport = DummyTransport(outlist)
|
||||
|
||||
xs.addObserver(xmlstream.STREAM_AUTHD_EVENT,
|
||||
self.authPassed)
|
||||
|
||||
# Go...
|
||||
xs.connectionMade()
|
||||
xs.dataReceived(b"<stream:stream xmlns='jabber:component:accept' xmlns:stream='http://etherx.jabber.org/streams' from='cjid' id='12345'>")
|
||||
|
||||
# Calculate what we expect the handshake value to be
|
||||
hv = sha1(b"12345" + b"secret").hexdigest().encode('ascii')
|
||||
|
||||
self.assertEqual(outlist[1], b"<handshake>" + hv + b"</handshake>")
|
||||
|
||||
xs.dataReceived("<handshake/>")
|
||||
|
||||
self.assertEqual(self.authComplete, True)
|
||||
|
||||
|
||||
|
||||
class ServiceTests(unittest.TestCase):
|
||||
"""
|
||||
Tests for L{component.Service}.
|
||||
"""
|
||||
|
||||
def test_interface(self):
|
||||
"""
|
||||
L{component.Service} implements L{ijabber.IService}.
|
||||
"""
|
||||
service = component.Service()
|
||||
verifyObject(ijabber.IService, service)
|
||||
|
||||
|
||||
|
||||
class JabberServiceHarness(component.Service):
|
||||
def __init__(self):
|
||||
self.componentConnectedFlag = False
|
||||
self.componentDisconnectedFlag = False
|
||||
self.transportConnectedFlag = False
|
||||
|
||||
def componentConnected(self, xmlstream):
|
||||
self.componentConnectedFlag = True
|
||||
|
||||
def componentDisconnected(self):
|
||||
self.componentDisconnectedFlag = True
|
||||
|
||||
def transportConnected(self, xmlstream):
|
||||
self.transportConnectedFlag = True
|
||||
|
||||
|
||||
class JabberServiceManagerTests(unittest.TestCase):
|
||||
def testSM(self):
|
||||
# Setup service manager and test harnes
|
||||
sm = component.ServiceManager("foo", "password")
|
||||
svc = JabberServiceHarness()
|
||||
svc.setServiceParent(sm)
|
||||
|
||||
# Create a write list
|
||||
wlist = []
|
||||
|
||||
# Setup a XmlStream
|
||||
xs = sm.getFactory().buildProtocol(None)
|
||||
xs.transport = self
|
||||
xs.transport.write = wlist.append
|
||||
|
||||
# Indicate that it's connected
|
||||
xs.connectionMade()
|
||||
|
||||
# Ensure the test service harness got notified
|
||||
self.assertEqual(True, svc.transportConnectedFlag)
|
||||
|
||||
# Jump ahead and pretend like the stream got auth'd
|
||||
xs.dispatch(xs, xmlstream.STREAM_AUTHD_EVENT)
|
||||
|
||||
# Ensure the test service harness got notified
|
||||
self.assertEqual(True, svc.componentConnectedFlag)
|
||||
|
||||
# Pretend to drop the connection
|
||||
xs.connectionLost(None)
|
||||
|
||||
# Ensure the test service harness got notified
|
||||
self.assertEqual(True, svc.componentDisconnectedFlag)
|
||||
|
||||
|
||||
|
||||
class RouterTests(unittest.TestCase):
|
||||
"""
|
||||
Tests for L{component.Router}.
|
||||
"""
|
||||
|
||||
def test_addRoute(self):
|
||||
"""
|
||||
Test route registration and routing on incoming stanzas.
|
||||
"""
|
||||
router = component.Router()
|
||||
routed = []
|
||||
router.route = lambda element: routed.append(element)
|
||||
|
||||
pipe = XmlPipe()
|
||||
router.addRoute('example.org', pipe.sink)
|
||||
self.assertEqual(1, len(router.routes))
|
||||
self.assertEqual(pipe.sink, router.routes['example.org'])
|
||||
|
||||
element = domish.Element(('testns', 'test'))
|
||||
pipe.source.send(element)
|
||||
self.assertEqual([element], routed)
|
||||
|
||||
|
||||
def test_route(self):
|
||||
"""
|
||||
Test routing of a message.
|
||||
"""
|
||||
component1 = XmlPipe()
|
||||
component2 = XmlPipe()
|
||||
router = component.Router()
|
||||
router.addRoute('component1.example.org', component1.sink)
|
||||
router.addRoute('component2.example.org', component2.sink)
|
||||
|
||||
outgoing = []
|
||||
component2.source.addObserver('/*',
|
||||
lambda element: outgoing.append(element))
|
||||
stanza = domish.Element((None, 'presence'))
|
||||
stanza['from'] = 'component1.example.org'
|
||||
stanza['to'] = 'component2.example.org'
|
||||
component1.source.send(stanza)
|
||||
self.assertEqual([stanza], outgoing)
|
||||
|
||||
|
||||
def test_routeDefault(self):
|
||||
"""
|
||||
Test routing of a message using the default route.
|
||||
|
||||
The default route is the one with L{None} as its key in the
|
||||
routing table. It is taken when there is no more specific route
|
||||
in the routing table that matches the stanza's destination.
|
||||
"""
|
||||
component1 = XmlPipe()
|
||||
s2s = XmlPipe()
|
||||
router = component.Router()
|
||||
router.addRoute('component1.example.org', component1.sink)
|
||||
router.addRoute(None, s2s.sink)
|
||||
|
||||
outgoing = []
|
||||
s2s.source.addObserver('/*', lambda element: outgoing.append(element))
|
||||
stanza = domish.Element((None, 'presence'))
|
||||
stanza['from'] = 'component1.example.org'
|
||||
stanza['to'] = 'example.com'
|
||||
component1.source.send(stanza)
|
||||
self.assertEqual([stanza], outgoing)
|
||||
|
||||
|
||||
|
||||
class ListenComponentAuthenticatorTests(unittest.TestCase):
|
||||
"""
|
||||
Tests for L{component.ListenComponentAuthenticator}.
|
||||
"""
|
||||
|
||||
def setUp(self):
|
||||
self.output = []
|
||||
authenticator = component.ListenComponentAuthenticator('secret')
|
||||
self.xmlstream = xmlstream.XmlStream(authenticator)
|
||||
self.xmlstream.send = self.output.append
|
||||
|
||||
|
||||
def loseConnection(self):
|
||||
"""
|
||||
Stub loseConnection because we are a transport.
|
||||
"""
|
||||
self.xmlstream.connectionLost("no reason")
|
||||
|
||||
|
||||
def test_streamStarted(self):
|
||||
"""
|
||||
The received stream header should set several attributes.
|
||||
"""
|
||||
observers = []
|
||||
|
||||
def addOnetimeObserver(event, observerfn):
|
||||
observers.append((event, observerfn))
|
||||
|
||||
xs = self.xmlstream
|
||||
xs.addOnetimeObserver = addOnetimeObserver
|
||||
|
||||
xs.makeConnection(self)
|
||||
self.assertIdentical(None, xs.sid)
|
||||
self.assertFalse(xs._headerSent)
|
||||
|
||||
xs.dataReceived("<stream:stream xmlns='jabber:component:accept' "
|
||||
"xmlns:stream='http://etherx.jabber.org/streams' "
|
||||
"to='component.example.org'>")
|
||||
self.assertEqual((0, 0), xs.version)
|
||||
self.assertNotIdentical(None, xs.sid)
|
||||
self.assertTrue(xs._headerSent)
|
||||
self.assertEqual(('/*', xs.authenticator.onElement), observers[-1])
|
||||
|
||||
|
||||
def test_streamStartedWrongNamespace(self):
|
||||
"""
|
||||
The received stream header should have a correct namespace.
|
||||
"""
|
||||
streamErrors = []
|
||||
|
||||
xs = self.xmlstream
|
||||
xs.sendStreamError = streamErrors.append
|
||||
xs.makeConnection(self)
|
||||
xs.dataReceived("<stream:stream xmlns='jabber:client' "
|
||||
"xmlns:stream='http://etherx.jabber.org/streams' "
|
||||
"to='component.example.org'>")
|
||||
self.assertEqual(1, len(streamErrors))
|
||||
self.assertEqual('invalid-namespace', streamErrors[-1].condition)
|
||||
|
||||
|
||||
def test_streamStartedNoTo(self):
|
||||
"""
|
||||
The received stream header should have a 'to' attribute.
|
||||
"""
|
||||
streamErrors = []
|
||||
|
||||
xs = self.xmlstream
|
||||
xs.sendStreamError = streamErrors.append
|
||||
xs.makeConnection(self)
|
||||
xs.dataReceived("<stream:stream xmlns='jabber:component:accept' "
|
||||
"xmlns:stream='http://etherx.jabber.org/streams'>")
|
||||
self.assertEqual(1, len(streamErrors))
|
||||
self.assertEqual('improper-addressing', streamErrors[-1].condition)
|
||||
|
||||
|
||||
def test_onElement(self):
|
||||
"""
|
||||
We expect a handshake element with a hash.
|
||||
"""
|
||||
handshakes = []
|
||||
|
||||
xs = self.xmlstream
|
||||
xs.authenticator.onHandshake = handshakes.append
|
||||
|
||||
handshake = domish.Element(('jabber:component:accept', 'handshake'))
|
||||
handshake.addContent(u'1234')
|
||||
xs.authenticator.onElement(handshake)
|
||||
self.assertEqual('1234', handshakes[-1])
|
||||
|
||||
def test_onElementNotHandshake(self):
|
||||
"""
|
||||
Reject elements that are not handshakes
|
||||
"""
|
||||
handshakes = []
|
||||
streamErrors = []
|
||||
|
||||
xs = self.xmlstream
|
||||
xs.authenticator.onHandshake = handshakes.append
|
||||
xs.sendStreamError = streamErrors.append
|
||||
|
||||
element = domish.Element(('jabber:component:accept', 'message'))
|
||||
xs.authenticator.onElement(element)
|
||||
self.assertFalse(handshakes)
|
||||
self.assertEqual('not-authorized', streamErrors[-1].condition)
|
||||
|
||||
|
||||
def test_onHandshake(self):
|
||||
"""
|
||||
Receiving a handshake matching the secret authenticates the stream.
|
||||
"""
|
||||
authd = []
|
||||
|
||||
def authenticated(xs):
|
||||
authd.append(xs)
|
||||
|
||||
xs = self.xmlstream
|
||||
xs.addOnetimeObserver(xmlstream.STREAM_AUTHD_EVENT, authenticated)
|
||||
xs.sid = u'1234'
|
||||
theHash = '32532c0f7dbf1253c095b18b18e36d38d94c1256'
|
||||
xs.authenticator.onHandshake(theHash)
|
||||
self.assertEqual('<handshake/>', self.output[-1])
|
||||
self.assertEqual(1, len(authd))
|
||||
|
||||
|
||||
def test_onHandshakeWrongHash(self):
|
||||
"""
|
||||
Receiving a bad handshake should yield a stream error.
|
||||
"""
|
||||
streamErrors = []
|
||||
authd = []
|
||||
|
||||
def authenticated(xs):
|
||||
authd.append(xs)
|
||||
|
||||
xs = self.xmlstream
|
||||
xs.addOnetimeObserver(xmlstream.STREAM_AUTHD_EVENT, authenticated)
|
||||
xs.sendStreamError = streamErrors.append
|
||||
|
||||
xs.sid = u'1234'
|
||||
theHash = '1234'
|
||||
xs.authenticator.onHandshake(theHash)
|
||||
self.assertEqual('not-authorized', streamErrors[-1].condition)
|
||||
self.assertEqual(0, len(authd))
|
||||
|
||||
|
||||
|
||||
class XMPPComponentServerFactoryTests(unittest.TestCase):
|
||||
"""
|
||||
Tests for L{component.XMPPComponentServerFactory}.
|
||||
"""
|
||||
|
||||
def setUp(self):
|
||||
self.router = component.Router()
|
||||
self.factory = component.XMPPComponentServerFactory(self.router,
|
||||
'secret')
|
||||
self.xmlstream = self.factory.buildProtocol(None)
|
||||
self.xmlstream.thisEntity = JID('component.example.org')
|
||||
|
||||
|
||||
def test_makeConnection(self):
|
||||
"""
|
||||
A new connection increases the stream serial count. No logs by default.
|
||||
"""
|
||||
self.xmlstream.dispatch(self.xmlstream,
|
||||
xmlstream.STREAM_CONNECTED_EVENT)
|
||||
self.assertEqual(0, self.xmlstream.serial)
|
||||
self.assertEqual(1, self.factory.serial)
|
||||
self.assertIdentical(None, self.xmlstream.rawDataInFn)
|
||||
self.assertIdentical(None, self.xmlstream.rawDataOutFn)
|
||||
|
||||
|
||||
def test_makeConnectionLogTraffic(self):
|
||||
"""
|
||||
Setting logTraffic should set up raw data loggers.
|
||||
"""
|
||||
self.factory.logTraffic = True
|
||||
self.xmlstream.dispatch(self.xmlstream,
|
||||
xmlstream.STREAM_CONNECTED_EVENT)
|
||||
self.assertNotIdentical(None, self.xmlstream.rawDataInFn)
|
||||
self.assertNotIdentical(None, self.xmlstream.rawDataOutFn)
|
||||
|
||||
|
||||
def test_onError(self):
|
||||
"""
|
||||
An observer for stream errors should trigger onError to log it.
|
||||
"""
|
||||
self.xmlstream.dispatch(self.xmlstream,
|
||||
xmlstream.STREAM_CONNECTED_EVENT)
|
||||
|
||||
class TestError(Exception):
|
||||
pass
|
||||
|
||||
reason = failure.Failure(TestError())
|
||||
self.xmlstream.dispatch(reason, xmlstream.STREAM_ERROR_EVENT)
|
||||
self.assertEqual(1, len(self.flushLoggedErrors(TestError)))
|
||||
|
||||
|
||||
def test_connectionInitialized(self):
|
||||
"""
|
||||
Make sure a new stream is added to the routing table.
|
||||
"""
|
||||
self.xmlstream.dispatch(self.xmlstream, xmlstream.STREAM_AUTHD_EVENT)
|
||||
self.assertIn('component.example.org', self.router.routes)
|
||||
self.assertIdentical(self.xmlstream,
|
||||
self.router.routes['component.example.org'])
|
||||
|
||||
|
||||
def test_connectionLost(self):
|
||||
"""
|
||||
Make sure a stream is removed from the routing table on disconnect.
|
||||
"""
|
||||
self.xmlstream.dispatch(self.xmlstream, xmlstream.STREAM_AUTHD_EVENT)
|
||||
self.xmlstream.dispatch(None, xmlstream.STREAM_END_EVENT)
|
||||
self.assertNotIn('component.example.org', self.router.routes)
|
||||
@@ -0,0 +1,36 @@
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
Tests for L{twisted.words.protocols.jabber.jstrports}.
|
||||
"""
|
||||
|
||||
from __future__ import absolute_import, division
|
||||
|
||||
from twisted.trial import unittest
|
||||
|
||||
from twisted.words.protocols.jabber import jstrports
|
||||
from twisted.application.internet import TCPClient
|
||||
|
||||
|
||||
class JabberStrPortsPlaceHolderTests(unittest.TestCase):
|
||||
"""
|
||||
Tests for L{jstrports}
|
||||
"""
|
||||
|
||||
def test_parse(self):
|
||||
"""
|
||||
L{jstrports.parse} accepts an endpoint description string and returns a
|
||||
tuple and dict of parsed endpoint arguments.
|
||||
"""
|
||||
expected = ('TCP', ('DOMAIN', 65535, 'Factory'), {})
|
||||
got = jstrports.parse("tcp:DOMAIN:65535", "Factory")
|
||||
self.assertEqual(expected, got)
|
||||
|
||||
|
||||
def test_client(self):
|
||||
"""
|
||||
L{jstrports.client} returns a L{TCPClient} service.
|
||||
"""
|
||||
got = jstrports.client("tcp:DOMAIN:65535", "Factory")
|
||||
self.assertIsInstance(got, TCPClient)
|
||||
@@ -0,0 +1,115 @@
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
from twisted.trial import unittest
|
||||
|
||||
from twisted.words.protocols.jabber.xmpp_stringprep import (
|
||||
nodeprep, resourceprep, nameprep)
|
||||
|
||||
|
||||
|
||||
class DeprecationTests(unittest.TestCase):
|
||||
"""
|
||||
Deprecations in L{twisted.words.protocols.jabber.xmpp_stringprep}.
|
||||
"""
|
||||
def test_crippled(self):
|
||||
"""
|
||||
L{xmpp_stringprep.crippled} is deprecated and always returns C{False}.
|
||||
"""
|
||||
from twisted.words.protocols.jabber.xmpp_stringprep import crippled
|
||||
warnings = self.flushWarnings(
|
||||
offendingFunctions=[self.test_crippled])
|
||||
self.assertEqual(DeprecationWarning, warnings[0]['category'])
|
||||
self.assertEqual(
|
||||
"twisted.words.protocols.jabber.xmpp_stringprep.crippled was "
|
||||
"deprecated in Twisted 13.1.0: crippled is always False",
|
||||
warnings[0]['message'])
|
||||
self.assertEqual(1, len(warnings))
|
||||
self.assertEqual(crippled, False)
|
||||
|
||||
|
||||
|
||||
class XMPPStringPrepTests(unittest.TestCase):
|
||||
"""
|
||||
The nodeprep stringprep profile is similar to the resourceprep profile,
|
||||
but does an extra mapping of characters (table B.2) and disallows
|
||||
more characters (table C.1.1 and eight extra punctuation characters).
|
||||
Due to this similarity, the resourceprep tests are more extensive, and
|
||||
the nodeprep tests only address the mappings additional restrictions.
|
||||
|
||||
The nameprep profile is nearly identical to the nameprep implementation in
|
||||
L{encodings.idna}, but that implementation assumes the C{UseSTD4ASCIIRules}
|
||||
flag to be false. This implementation assumes it to be true, and restricts
|
||||
the allowed set of characters. The tests here only check for the
|
||||
differences.
|
||||
"""
|
||||
|
||||
def testResourcePrep(self):
|
||||
self.assertEqual(resourceprep.prepare(u'resource'), u'resource')
|
||||
self.assertNotEqual(resourceprep.prepare(u'Resource'), u'resource')
|
||||
self.assertEqual(resourceprep.prepare(u' '), u' ')
|
||||
|
||||
self.assertEqual(resourceprep.prepare(u'Henry \u2163'), u'Henry IV')
|
||||
self.assertEqual(resourceprep.prepare(u'foo\xad\u034f\u1806\u180b'
|
||||
u'bar\u200b\u2060'
|
||||
u'baz\ufe00\ufe08\ufe0f\ufeff'),
|
||||
u'foobarbaz')
|
||||
self.assertEqual(resourceprep.prepare(u'\u00a0'), u' ')
|
||||
self.assertRaises(UnicodeError, resourceprep.prepare, u'\u1680')
|
||||
self.assertEqual(resourceprep.prepare(u'\u2000'), u' ')
|
||||
self.assertEqual(resourceprep.prepare(u'\u200b'), u'')
|
||||
self.assertRaises(UnicodeError, resourceprep.prepare, u'\u0010\u007f')
|
||||
self.assertRaises(UnicodeError, resourceprep.prepare, u'\u0085')
|
||||
self.assertRaises(UnicodeError, resourceprep.prepare, u'\u180e')
|
||||
self.assertEqual(resourceprep.prepare(u'\ufeff'), u'')
|
||||
self.assertRaises(UnicodeError, resourceprep.prepare, u'\uf123')
|
||||
self.assertRaises(UnicodeError, resourceprep.prepare, u'\U000f1234')
|
||||
self.assertRaises(UnicodeError, resourceprep.prepare, u'\U0010f234')
|
||||
self.assertRaises(UnicodeError, resourceprep.prepare, u'\U0008fffe')
|
||||
self.assertRaises(UnicodeError, resourceprep.prepare, u'\U0010ffff')
|
||||
self.assertRaises(UnicodeError, resourceprep.prepare, u'\udf42')
|
||||
self.assertRaises(UnicodeError, resourceprep.prepare, u'\ufffd')
|
||||
self.assertRaises(UnicodeError, resourceprep.prepare, u'\u2ff5')
|
||||
self.assertEqual(resourceprep.prepare(u'\u0341'), u'\u0301')
|
||||
self.assertRaises(UnicodeError, resourceprep.prepare, u'\u200e')
|
||||
self.assertRaises(UnicodeError, resourceprep.prepare, u'\u202a')
|
||||
self.assertRaises(UnicodeError, resourceprep.prepare, u'\U000e0001')
|
||||
self.assertRaises(UnicodeError, resourceprep.prepare, u'\U000e0042')
|
||||
self.assertRaises(UnicodeError, resourceprep.prepare, u'foo\u05bebar')
|
||||
self.assertRaises(UnicodeError, resourceprep.prepare, u'foo\ufd50bar')
|
||||
#self.assertEqual(resourceprep.prepare(u'foo\ufb38bar'),
|
||||
# u'foo\u064ebar')
|
||||
self.assertRaises(UnicodeError, resourceprep.prepare, u'\u06271')
|
||||
self.assertEqual(resourceprep.prepare(u'\u06271\u0628'),
|
||||
u'\u06271\u0628')
|
||||
self.assertRaises(UnicodeError, resourceprep.prepare, u'\U000e0002')
|
||||
|
||||
|
||||
def testNodePrep(self):
|
||||
self.assertEqual(nodeprep.prepare(u'user'), u'user')
|
||||
self.assertEqual(nodeprep.prepare(u'User'), u'user')
|
||||
self.assertRaises(UnicodeError, nodeprep.prepare, u'us&er')
|
||||
|
||||
|
||||
def test_nodeprepUnassignedInUnicode32(self):
|
||||
"""
|
||||
Make sure unassigned code points from Unicode 3.2 are rejected.
|
||||
"""
|
||||
self.assertRaises(UnicodeError, nodeprep.prepare, u'\u1d39')
|
||||
|
||||
|
||||
def testNamePrep(self):
|
||||
self.assertEqual(nameprep.prepare(u'example.com'), u'example.com')
|
||||
self.assertEqual(nameprep.prepare(u'Example.com'), u'example.com')
|
||||
self.assertRaises(UnicodeError, nameprep.prepare, u'ex@mple.com')
|
||||
self.assertRaises(UnicodeError, nameprep.prepare, u'-example.com')
|
||||
self.assertRaises(UnicodeError, nameprep.prepare, u'example-.com')
|
||||
|
||||
self.assertEqual(nameprep.prepare(u'stra\u00dfe.example.com'),
|
||||
u'strasse.example.com')
|
||||
|
||||
def test_nameprepTrailingDot(self):
|
||||
"""
|
||||
A trailing dot in domain names is preserved.
|
||||
"""
|
||||
self.assertEqual(nameprep.prepare(u'example.com.'), u'example.com.')
|
||||
@@ -0,0 +1,78 @@
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
from twisted.cred import credentials, error
|
||||
from twisted.words import tap
|
||||
from twisted.trial import unittest
|
||||
|
||||
|
||||
|
||||
class WordsTapTests(unittest.TestCase):
|
||||
"""
|
||||
Ensures that the twisted.words.tap API works.
|
||||
"""
|
||||
|
||||
PASSWD_TEXT = b"admin:admin\njoe:foo\n"
|
||||
admin = credentials.UsernamePassword(b'admin', b'admin')
|
||||
joeWrong = credentials.UsernamePassword(b'joe', b'bar')
|
||||
|
||||
|
||||
def setUp(self):
|
||||
"""
|
||||
Create a file with two users.
|
||||
"""
|
||||
self.filename = self.mktemp()
|
||||
self.file = open(self.filename, 'wb')
|
||||
self.file.write(self.PASSWD_TEXT)
|
||||
self.file.flush()
|
||||
|
||||
|
||||
def tearDown(self):
|
||||
"""
|
||||
Close the dummy user database.
|
||||
"""
|
||||
self.file.close()
|
||||
|
||||
|
||||
def test_hostname(self):
|
||||
"""
|
||||
Tests that the --hostname parameter gets passed to Options.
|
||||
"""
|
||||
opt = tap.Options()
|
||||
opt.parseOptions(['--hostname', 'myhost'])
|
||||
self.assertEqual(opt['hostname'], 'myhost')
|
||||
|
||||
|
||||
def test_passwd(self):
|
||||
"""
|
||||
Tests the --passwd command for backwards-compatibility.
|
||||
"""
|
||||
opt = tap.Options()
|
||||
opt.parseOptions(['--passwd', self.file.name])
|
||||
self._loginTest(opt)
|
||||
|
||||
|
||||
def test_auth(self):
|
||||
"""
|
||||
Tests that the --auth command generates a checker.
|
||||
"""
|
||||
opt = tap.Options()
|
||||
opt.parseOptions(['--auth', 'file:'+self.file.name])
|
||||
self._loginTest(opt)
|
||||
|
||||
|
||||
def _loginTest(self, opt):
|
||||
"""
|
||||
This method executes both positive and negative authentication
|
||||
tests against whatever credentials checker has been stored in
|
||||
the Options class.
|
||||
|
||||
@param opt: An instance of L{tap.Options}.
|
||||
"""
|
||||
self.assertEqual(len(opt['credCheckers']), 1)
|
||||
checker = opt['credCheckers'][0]
|
||||
self.assertFailure(checker.requestAvatarId(self.joeWrong),
|
||||
error.UnauthorizedLogin)
|
||||
def _gotAvatar(username):
|
||||
self.assertEqual(username, self.admin.username)
|
||||
return checker.requestAvatarId(self.admin).addCallback(_gotAvatar)
|
||||
@@ -0,0 +1,348 @@
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
Test cases for twisted.words.xish.utility
|
||||
"""
|
||||
|
||||
from __future__ import absolute_import, division
|
||||
|
||||
from collections import OrderedDict
|
||||
|
||||
from twisted.trial import unittest
|
||||
|
||||
from twisted.words.xish import utility
|
||||
from twisted.words.xish.domish import Element
|
||||
from twisted.words.xish.utility import EventDispatcher
|
||||
|
||||
class CallbackTracker:
|
||||
"""
|
||||
Test helper for tracking callbacks.
|
||||
|
||||
Increases a counter on each call to L{call} and stores the object
|
||||
passed in the call.
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
self.called = 0
|
||||
self.obj = None
|
||||
|
||||
|
||||
def call(self, obj):
|
||||
self.called = self.called + 1
|
||||
self.obj = obj
|
||||
|
||||
|
||||
|
||||
class OrderedCallbackTracker:
|
||||
"""
|
||||
Test helper for tracking callbacks and their order.
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
self.callList = []
|
||||
|
||||
|
||||
def call1(self, object):
|
||||
self.callList.append(self.call1)
|
||||
|
||||
|
||||
def call2(self, object):
|
||||
self.callList.append(self.call2)
|
||||
|
||||
|
||||
def call3(self, object):
|
||||
self.callList.append(self.call3)
|
||||
|
||||
|
||||
|
||||
class EventDispatcherTests(unittest.TestCase):
|
||||
"""
|
||||
Tests for L{EventDispatcher}.
|
||||
"""
|
||||
|
||||
def testStuff(self):
|
||||
d = EventDispatcher()
|
||||
cb1 = CallbackTracker()
|
||||
cb2 = CallbackTracker()
|
||||
cb3 = CallbackTracker()
|
||||
|
||||
d.addObserver("/message/body", cb1.call)
|
||||
d.addObserver("/message", cb1.call)
|
||||
d.addObserver("/presence", cb2.call)
|
||||
d.addObserver("//event/testevent", cb3.call)
|
||||
|
||||
msg = Element(("ns", "message"))
|
||||
msg.addElement("body")
|
||||
|
||||
pres = Element(("ns", "presence"))
|
||||
pres.addElement("presence")
|
||||
|
||||
d.dispatch(msg)
|
||||
self.assertEqual(cb1.called, 2)
|
||||
self.assertEqual(cb1.obj, msg)
|
||||
self.assertEqual(cb2.called, 0)
|
||||
|
||||
d.dispatch(pres)
|
||||
self.assertEqual(cb1.called, 2)
|
||||
self.assertEqual(cb2.called, 1)
|
||||
self.assertEqual(cb2.obj, pres)
|
||||
self.assertEqual(cb3.called, 0)
|
||||
|
||||
d.dispatch(d, "//event/testevent")
|
||||
self.assertEqual(cb3.called, 1)
|
||||
self.assertEqual(cb3.obj, d)
|
||||
|
||||
d.removeObserver("/presence", cb2.call)
|
||||
d.dispatch(pres)
|
||||
self.assertEqual(cb2.called, 1)
|
||||
|
||||
|
||||
def test_addObserverTwice(self):
|
||||
"""
|
||||
Test adding two observers for the same query.
|
||||
|
||||
When the event is dispatched both of the observers need to be called.
|
||||
"""
|
||||
d = EventDispatcher()
|
||||
cb1 = CallbackTracker()
|
||||
cb2 = CallbackTracker()
|
||||
|
||||
d.addObserver("//event/testevent", cb1.call)
|
||||
d.addObserver("//event/testevent", cb2.call)
|
||||
d.dispatch(d, "//event/testevent")
|
||||
|
||||
self.assertEqual(cb1.called, 1)
|
||||
self.assertEqual(cb1.obj, d)
|
||||
self.assertEqual(cb2.called, 1)
|
||||
self.assertEqual(cb2.obj, d)
|
||||
|
||||
|
||||
def test_addObserverInDispatch(self):
|
||||
"""
|
||||
Test for registration of an observer during dispatch.
|
||||
"""
|
||||
d = EventDispatcher()
|
||||
msg = Element(("ns", "message"))
|
||||
cb = CallbackTracker()
|
||||
|
||||
def onMessage(_):
|
||||
d.addObserver("/message", cb.call)
|
||||
|
||||
d.addOnetimeObserver("/message", onMessage)
|
||||
|
||||
d.dispatch(msg)
|
||||
self.assertEqual(cb.called, 0)
|
||||
|
||||
d.dispatch(msg)
|
||||
self.assertEqual(cb.called, 1)
|
||||
|
||||
d.dispatch(msg)
|
||||
self.assertEqual(cb.called, 2)
|
||||
|
||||
|
||||
def test_addOnetimeObserverInDispatch(self):
|
||||
"""
|
||||
Test for registration of a onetime observer during dispatch.
|
||||
"""
|
||||
d = EventDispatcher()
|
||||
msg = Element(("ns", "message"))
|
||||
cb = CallbackTracker()
|
||||
|
||||
def onMessage(msg):
|
||||
d.addOnetimeObserver("/message", cb.call)
|
||||
|
||||
d.addOnetimeObserver("/message", onMessage)
|
||||
|
||||
d.dispatch(msg)
|
||||
self.assertEqual(cb.called, 0)
|
||||
|
||||
d.dispatch(msg)
|
||||
self.assertEqual(cb.called, 1)
|
||||
|
||||
d.dispatch(msg)
|
||||
self.assertEqual(cb.called, 1)
|
||||
|
||||
|
||||
def testOnetimeDispatch(self):
|
||||
d = EventDispatcher()
|
||||
msg = Element(("ns", "message"))
|
||||
cb = CallbackTracker()
|
||||
|
||||
d.addOnetimeObserver("/message", cb.call)
|
||||
d.dispatch(msg)
|
||||
self.assertEqual(cb.called, 1)
|
||||
d.dispatch(msg)
|
||||
self.assertEqual(cb.called, 1)
|
||||
|
||||
|
||||
def testDispatcherResult(self):
|
||||
d = EventDispatcher()
|
||||
msg = Element(("ns", "message"))
|
||||
pres = Element(("ns", "presence"))
|
||||
cb = CallbackTracker()
|
||||
|
||||
d.addObserver("/presence", cb.call)
|
||||
result = d.dispatch(msg)
|
||||
self.assertEqual(False, result)
|
||||
|
||||
result = d.dispatch(pres)
|
||||
self.assertEqual(True, result)
|
||||
|
||||
|
||||
def testOrderedXPathDispatch(self):
|
||||
d = EventDispatcher()
|
||||
cb = OrderedCallbackTracker()
|
||||
d.addObserver("/message/body", cb.call2)
|
||||
d.addObserver("/message", cb.call3, -1)
|
||||
d.addObserver("/message/body", cb.call1, 1)
|
||||
|
||||
msg = Element(("ns", "message"))
|
||||
msg.addElement("body")
|
||||
d.dispatch(msg)
|
||||
self.assertEqual(cb.callList, [cb.call1, cb.call2, cb.call3],
|
||||
"Calls out of order: %s" %
|
||||
repr([c.__name__ for c in cb.callList]))
|
||||
|
||||
|
||||
# Observers are put into CallbackLists that are then put into dictionaries
|
||||
# keyed by the event trigger. Upon removal of the last observer for a
|
||||
# particular event trigger, the (now empty) CallbackList and corresponding
|
||||
# event trigger should be removed from those dictionaries to prevent
|
||||
# slowdown and memory leakage.
|
||||
|
||||
def test_cleanUpRemoveEventObserver(self):
|
||||
"""
|
||||
Test observer clean-up after removeObserver for named events.
|
||||
"""
|
||||
|
||||
d = EventDispatcher()
|
||||
cb = CallbackTracker()
|
||||
|
||||
d.addObserver('//event/test', cb.call)
|
||||
d.dispatch(None, '//event/test')
|
||||
self.assertEqual(1, cb.called)
|
||||
d.removeObserver('//event/test', cb.call)
|
||||
self.assertEqual(0, len(d._eventObservers.pop(0)))
|
||||
|
||||
|
||||
def test_cleanUpRemoveXPathObserver(self):
|
||||
"""
|
||||
Test observer clean-up after removeObserver for XPath events.
|
||||
"""
|
||||
|
||||
d = EventDispatcher()
|
||||
cb = CallbackTracker()
|
||||
msg = Element((None, "message"))
|
||||
|
||||
d.addObserver('/message', cb.call)
|
||||
d.dispatch(msg)
|
||||
self.assertEqual(1, cb.called)
|
||||
d.removeObserver('/message', cb.call)
|
||||
self.assertEqual(0, len(d._xpathObservers.pop(0)))
|
||||
|
||||
|
||||
def test_cleanUpOnetimeEventObserver(self):
|
||||
"""
|
||||
Test observer clean-up after onetime named events.
|
||||
"""
|
||||
|
||||
d = EventDispatcher()
|
||||
cb = CallbackTracker()
|
||||
|
||||
d.addOnetimeObserver('//event/test', cb.call)
|
||||
d.dispatch(None, '//event/test')
|
||||
self.assertEqual(1, cb.called)
|
||||
self.assertEqual(0, len(d._eventObservers.pop(0)))
|
||||
|
||||
|
||||
def test_cleanUpOnetimeXPathObserver(self):
|
||||
"""
|
||||
Test observer clean-up after onetime XPath events.
|
||||
"""
|
||||
|
||||
d = EventDispatcher()
|
||||
cb = CallbackTracker()
|
||||
msg = Element((None, "message"))
|
||||
|
||||
d.addOnetimeObserver('/message', cb.call)
|
||||
d.dispatch(msg)
|
||||
self.assertEqual(1, cb.called)
|
||||
self.assertEqual(0, len(d._xpathObservers.pop(0)))
|
||||
|
||||
|
||||
def test_observerRaisingException(self):
|
||||
"""
|
||||
Test that exceptions in observers do not bubble up to dispatch.
|
||||
|
||||
The exceptions raised in observers should be logged and other
|
||||
observers should be called as if nothing happened.
|
||||
"""
|
||||
|
||||
class OrderedCallbackList(utility.CallbackList):
|
||||
def __init__(self):
|
||||
self.callbacks = OrderedDict()
|
||||
|
||||
class TestError(Exception):
|
||||
pass
|
||||
|
||||
def raiseError(_):
|
||||
raise TestError()
|
||||
|
||||
d = EventDispatcher()
|
||||
cb = CallbackTracker()
|
||||
|
||||
originalCallbackList = utility.CallbackList
|
||||
|
||||
try:
|
||||
utility.CallbackList = OrderedCallbackList
|
||||
|
||||
d.addObserver('//event/test', raiseError)
|
||||
d.addObserver('//event/test', cb.call)
|
||||
try:
|
||||
d.dispatch(None, '//event/test')
|
||||
except TestError:
|
||||
self.fail("TestError raised. Should have been logged instead.")
|
||||
|
||||
self.assertEqual(1, len(self.flushLoggedErrors(TestError)))
|
||||
self.assertEqual(1, cb.called)
|
||||
finally:
|
||||
utility.CallbackList = originalCallbackList
|
||||
|
||||
|
||||
|
||||
class XmlPipeTests(unittest.TestCase):
|
||||
"""
|
||||
Tests for L{twisted.words.xish.utility.XmlPipe}.
|
||||
"""
|
||||
|
||||
def setUp(self):
|
||||
self.pipe = utility.XmlPipe()
|
||||
|
||||
|
||||
def test_sendFromSource(self):
|
||||
"""
|
||||
Send an element from the source and observe it from the sink.
|
||||
"""
|
||||
def cb(obj):
|
||||
called.append(obj)
|
||||
|
||||
called = []
|
||||
self.pipe.sink.addObserver('/test[@xmlns="testns"]', cb)
|
||||
element = Element(('testns', 'test'))
|
||||
self.pipe.source.send(element)
|
||||
self.assertEqual([element], called)
|
||||
|
||||
|
||||
def test_sendFromSink(self):
|
||||
"""
|
||||
Send an element from the sink and observe it from the source.
|
||||
"""
|
||||
def cb(obj):
|
||||
called.append(obj)
|
||||
|
||||
called = []
|
||||
self.pipe.source.addObserver('/test[@xmlns="testns"]', cb)
|
||||
element = Element(('testns', 'test'))
|
||||
self.pipe.sink.send(element)
|
||||
self.assertEqual([element], called)
|
||||
@@ -0,0 +1,84 @@
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
Tests for L{twisted.words.xmpproutertap}.
|
||||
"""
|
||||
|
||||
from twisted.application import internet
|
||||
from twisted.trial import unittest
|
||||
from twisted.words import xmpproutertap as tap
|
||||
from twisted.words.protocols.jabber import component
|
||||
|
||||
class XMPPRouterTapTests(unittest.TestCase):
|
||||
|
||||
def test_port(self):
|
||||
"""
|
||||
The port option is recognised as a parameter.
|
||||
"""
|
||||
opt = tap.Options()
|
||||
opt.parseOptions(['--port', '7001'])
|
||||
self.assertEqual(opt['port'], '7001')
|
||||
|
||||
|
||||
def test_portDefault(self):
|
||||
"""
|
||||
The port option has '5347' as default value
|
||||
"""
|
||||
opt = tap.Options()
|
||||
opt.parseOptions([])
|
||||
self.assertEqual(opt['port'], 'tcp:5347:interface=127.0.0.1')
|
||||
|
||||
|
||||
def test_secret(self):
|
||||
"""
|
||||
The secret option is recognised as a parameter.
|
||||
"""
|
||||
opt = tap.Options()
|
||||
opt.parseOptions(['--secret', 'hushhush'])
|
||||
self.assertEqual(opt['secret'], 'hushhush')
|
||||
|
||||
|
||||
def test_secretDefault(self):
|
||||
"""
|
||||
The secret option has 'secret' as default value
|
||||
"""
|
||||
opt = tap.Options()
|
||||
opt.parseOptions([])
|
||||
self.assertEqual(opt['secret'], 'secret')
|
||||
|
||||
|
||||
def test_verbose(self):
|
||||
"""
|
||||
The verbose option is recognised as a flag.
|
||||
"""
|
||||
opt = tap.Options()
|
||||
opt.parseOptions(['--verbose'])
|
||||
self.assertTrue(opt['verbose'])
|
||||
|
||||
|
||||
def test_makeService(self):
|
||||
"""
|
||||
The service gets set up with a router and factory.
|
||||
"""
|
||||
opt = tap.Options()
|
||||
opt.parseOptions([])
|
||||
s = tap.makeService(opt)
|
||||
self.assertIsInstance(s, internet.StreamServerEndpointService)
|
||||
self.assertEqual('127.0.0.1', s.endpoint._interface)
|
||||
self.assertEqual(5347, s.endpoint._port)
|
||||
factory = s.factory
|
||||
self.assertIsInstance(factory, component.XMPPComponentServerFactory)
|
||||
self.assertIsInstance(factory.router, component.Router)
|
||||
self.assertEqual('secret', factory.secret)
|
||||
self.assertFalse(factory.logTraffic)
|
||||
|
||||
|
||||
def test_makeServiceVerbose(self):
|
||||
"""
|
||||
The verbose flag enables traffic logging.
|
||||
"""
|
||||
opt = tap.Options()
|
||||
opt.parseOptions(['--verbose'])
|
||||
s = tap.makeService(opt)
|
||||
self.assertTrue(s.factory.logTraffic)
|
||||
@@ -0,0 +1,10 @@
|
||||
# -*- test-case-name: twisted.words.test -*-
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
|
||||
"""
|
||||
|
||||
Twisted X-ish: XML-ish DOM and XPath-ish engine
|
||||
|
||||
"""
|
||||
@@ -0,0 +1,337 @@
|
||||
# -*- test-case-name: twisted.words.test.test_xpath -*-
|
||||
#
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
XPath query support.
|
||||
|
||||
This module provides L{XPathQuery} to match
|
||||
L{domish.Element<twisted.words.xish.domish.Element>} instances against
|
||||
XPath-like expressions.
|
||||
"""
|
||||
|
||||
from __future__ import absolute_import, division
|
||||
|
||||
from io import StringIO
|
||||
|
||||
from twisted.python.compat import StringType, unicode
|
||||
|
||||
class LiteralValue(unicode):
|
||||
def value(self, elem):
|
||||
return self
|
||||
|
||||
|
||||
class IndexValue:
|
||||
def __init__(self, index):
|
||||
self.index = int(index) - 1
|
||||
|
||||
def value(self, elem):
|
||||
return elem.children[self.index]
|
||||
|
||||
|
||||
class AttribValue:
|
||||
def __init__(self, attribname):
|
||||
self.attribname = attribname
|
||||
if self.attribname == "xmlns":
|
||||
self.value = self.value_ns
|
||||
|
||||
def value_ns(self, elem):
|
||||
return elem.uri
|
||||
|
||||
def value(self, elem):
|
||||
if self.attribname in elem.attributes:
|
||||
return elem.attributes[self.attribname]
|
||||
else:
|
||||
return None
|
||||
|
||||
|
||||
class CompareValue:
|
||||
def __init__(self, lhs, op, rhs):
|
||||
self.lhs = lhs
|
||||
self.rhs = rhs
|
||||
if op == "=":
|
||||
self.value = self._compareEqual
|
||||
else:
|
||||
self.value = self._compareNotEqual
|
||||
|
||||
def _compareEqual(self, elem):
|
||||
return self.lhs.value(elem) == self.rhs.value(elem)
|
||||
|
||||
def _compareNotEqual(self, elem):
|
||||
return self.lhs.value(elem) != self.rhs.value(elem)
|
||||
|
||||
|
||||
class BooleanValue:
|
||||
"""
|
||||
Provide boolean XPath expression operators.
|
||||
|
||||
@ivar lhs: Left hand side expression of the operator.
|
||||
@ivar op: The operator. One of C{'and'}, C{'or'}.
|
||||
@ivar rhs: Right hand side expression of the operator.
|
||||
@ivar value: Reference to the method that will calculate the value of
|
||||
this expression given an element.
|
||||
"""
|
||||
def __init__(self, lhs, op, rhs):
|
||||
self.lhs = lhs
|
||||
self.rhs = rhs
|
||||
if op == "and":
|
||||
self.value = self._booleanAnd
|
||||
else:
|
||||
self.value = self._booleanOr
|
||||
|
||||
def _booleanAnd(self, elem):
|
||||
"""
|
||||
Calculate boolean and of the given expressions given an element.
|
||||
|
||||
@param elem: The element to calculate the value of the expression from.
|
||||
"""
|
||||
return self.lhs.value(elem) and self.rhs.value(elem)
|
||||
|
||||
def _booleanOr(self, elem):
|
||||
"""
|
||||
Calculate boolean or of the given expressions given an element.
|
||||
|
||||
@param elem: The element to calculate the value of the expression from.
|
||||
"""
|
||||
return self.lhs.value(elem) or self.rhs.value(elem)
|
||||
|
||||
|
||||
def Function(fname):
|
||||
"""
|
||||
Internal method which selects the function object
|
||||
"""
|
||||
klassname = "_%s_Function" % fname
|
||||
c = globals()[klassname]()
|
||||
return c
|
||||
|
||||
|
||||
class _not_Function:
|
||||
def __init__(self):
|
||||
self.baseValue = None
|
||||
|
||||
def setParams(self, baseValue):
|
||||
self.baseValue = baseValue
|
||||
|
||||
def value(self, elem):
|
||||
return not self.baseValue.value(elem)
|
||||
|
||||
|
||||
class _text_Function:
|
||||
def setParams(self):
|
||||
pass
|
||||
|
||||
def value(self, elem):
|
||||
return unicode(elem)
|
||||
|
||||
|
||||
class _Location:
|
||||
def __init__(self):
|
||||
self.predicates = []
|
||||
self.elementName = None
|
||||
self.childLocation = None
|
||||
|
||||
def matchesPredicates(self, elem):
|
||||
if self.elementName != None and self.elementName != elem.name:
|
||||
return 0
|
||||
|
||||
for p in self.predicates:
|
||||
if not p.value(elem):
|
||||
return 0
|
||||
|
||||
return 1
|
||||
|
||||
def matches(self, elem):
|
||||
if not self.matchesPredicates(elem):
|
||||
return 0
|
||||
|
||||
if self.childLocation != None:
|
||||
for c in elem.elements():
|
||||
if self.childLocation.matches(c):
|
||||
return 1
|
||||
else:
|
||||
return 1
|
||||
|
||||
return 0
|
||||
|
||||
def queryForString(self, elem, resultbuf):
|
||||
if not self.matchesPredicates(elem):
|
||||
return
|
||||
|
||||
if self.childLocation != None:
|
||||
for c in elem.elements():
|
||||
self.childLocation.queryForString(c, resultbuf)
|
||||
else:
|
||||
resultbuf.write(unicode(elem))
|
||||
|
||||
def queryForNodes(self, elem, resultlist):
|
||||
if not self.matchesPredicates(elem):
|
||||
return
|
||||
|
||||
if self.childLocation != None:
|
||||
for c in elem.elements():
|
||||
self.childLocation.queryForNodes(c, resultlist)
|
||||
else:
|
||||
resultlist.append(elem)
|
||||
|
||||
def queryForStringList(self, elem, resultlist):
|
||||
if not self.matchesPredicates(elem):
|
||||
return
|
||||
|
||||
if self.childLocation != None:
|
||||
for c in elem.elements():
|
||||
self.childLocation.queryForStringList(c, resultlist)
|
||||
else:
|
||||
for c in elem.children:
|
||||
if isinstance(c, StringType):
|
||||
resultlist.append(c)
|
||||
|
||||
|
||||
class _AnyLocation:
|
||||
def __init__(self):
|
||||
self.predicates = []
|
||||
self.elementName = None
|
||||
self.childLocation = None
|
||||
|
||||
def matchesPredicates(self, elem):
|
||||
for p in self.predicates:
|
||||
if not p.value(elem):
|
||||
return 0
|
||||
return 1
|
||||
|
||||
def listParents(self, elem, parentlist):
|
||||
if elem.parent != None:
|
||||
self.listParents(elem.parent, parentlist)
|
||||
parentlist.append(elem.name)
|
||||
|
||||
def isRootMatch(self, elem):
|
||||
if (self.elementName == None or self.elementName == elem.name) and \
|
||||
self.matchesPredicates(elem):
|
||||
if self.childLocation != None:
|
||||
for c in elem.elements():
|
||||
if self.childLocation.matches(c):
|
||||
return True
|
||||
else:
|
||||
return True
|
||||
return False
|
||||
|
||||
def findFirstRootMatch(self, elem):
|
||||
if (self.elementName == None or self.elementName == elem.name) and \
|
||||
self.matchesPredicates(elem):
|
||||
# Thus far, the name matches and the predicates match,
|
||||
# now check into the children and find the first one
|
||||
# that matches the rest of the structure
|
||||
# the rest of the structure
|
||||
if self.childLocation != None:
|
||||
for c in elem.elements():
|
||||
if self.childLocation.matches(c):
|
||||
return c
|
||||
return None
|
||||
else:
|
||||
# No children locations; this is a match!
|
||||
return elem
|
||||
else:
|
||||
# Ok, predicates or name didn't match, so we need to start
|
||||
# down each child and treat it as the root and try
|
||||
# again
|
||||
for c in elem.elements():
|
||||
if self.matches(c):
|
||||
return c
|
||||
# No children matched...
|
||||
return None
|
||||
|
||||
def matches(self, elem):
|
||||
if self.isRootMatch(elem):
|
||||
return True
|
||||
else:
|
||||
# Ok, initial element isn't an exact match, walk
|
||||
# down each child and treat it as the root and try
|
||||
# again
|
||||
for c in elem.elements():
|
||||
if self.matches(c):
|
||||
return True
|
||||
# No children matched...
|
||||
return False
|
||||
|
||||
def queryForString(self, elem, resultbuf):
|
||||
raise NotImplementedError(
|
||||
"queryForString is not implemented for any location")
|
||||
|
||||
def queryForNodes(self, elem, resultlist):
|
||||
# First check to see if _this_ element is a root
|
||||
if self.isRootMatch(elem):
|
||||
resultlist.append(elem)
|
||||
|
||||
# Now check each child
|
||||
for c in elem.elements():
|
||||
self.queryForNodes(c, resultlist)
|
||||
|
||||
|
||||
def queryForStringList(self, elem, resultlist):
|
||||
if self.isRootMatch(elem):
|
||||
for c in elem.children:
|
||||
if isinstance(c, StringType):
|
||||
resultlist.append(c)
|
||||
for c in elem.elements():
|
||||
self.queryForStringList(c, resultlist)
|
||||
|
||||
|
||||
class XPathQuery:
|
||||
def __init__(self, queryStr):
|
||||
self.queryStr = queryStr
|
||||
# Prevent a circular import issue, as xpathparser imports this module.
|
||||
from twisted.words.xish.xpathparser import (XPathParser,
|
||||
XPathParserScanner)
|
||||
parser = XPathParser(XPathParserScanner(queryStr))
|
||||
self.baseLocation = getattr(parser, 'XPATH')()
|
||||
|
||||
def __hash__(self):
|
||||
return self.queryStr.__hash__()
|
||||
|
||||
def matches(self, elem):
|
||||
return self.baseLocation.matches(elem)
|
||||
|
||||
def queryForString(self, elem):
|
||||
result = StringIO()
|
||||
self.baseLocation.queryForString(elem, result)
|
||||
return result.getvalue()
|
||||
|
||||
def queryForNodes(self, elem):
|
||||
result = []
|
||||
self.baseLocation.queryForNodes(elem, result)
|
||||
if len(result) == 0:
|
||||
return None
|
||||
else:
|
||||
return result
|
||||
|
||||
def queryForStringList(self, elem):
|
||||
result = []
|
||||
self.baseLocation.queryForStringList(elem, result)
|
||||
if len(result) == 0:
|
||||
return None
|
||||
else:
|
||||
return result
|
||||
|
||||
|
||||
__internedQueries = {}
|
||||
|
||||
def internQuery(queryString):
|
||||
if queryString not in __internedQueries:
|
||||
__internedQueries[queryString] = XPathQuery(queryString)
|
||||
return __internedQueries[queryString]
|
||||
|
||||
|
||||
def matches(xpathstr, elem):
|
||||
return internQuery(xpathstr).matches(elem)
|
||||
|
||||
|
||||
def queryForStringList(xpathstr, elem):
|
||||
return internQuery(xpathstr).queryForStringList(elem)
|
||||
|
||||
|
||||
def queryForString(xpathstr, elem):
|
||||
return internQuery(xpathstr).queryForString(elem)
|
||||
|
||||
|
||||
def queryForNodes(xpathstr, elem):
|
||||
return internQuery(xpathstr).queryForNodes(elem)
|
||||
@@ -0,0 +1,524 @@
|
||||
# -*- test-case-name: twisted.words.test.test_xpath -*-
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
# pylint: disable=W9401,W9402
|
||||
|
||||
# DO NOT EDIT xpathparser.py!
|
||||
#
|
||||
# It is generated from xpathparser.g using Yapps. Make needed changes there.
|
||||
# This also means that the generated Python may not conform to Twisted's coding
|
||||
# standards, so it is wrapped in exec to prevent automated checkers from
|
||||
# complaining.
|
||||
|
||||
# HOWTO Generate me:
|
||||
#
|
||||
# 1.) Grab a copy of yapps2:
|
||||
# https://github.com/smurfix/yapps
|
||||
#
|
||||
# Note: Do NOT use the package in debian/ubuntu as it has incompatible
|
||||
# modifications. The original at http://theory.stanford.edu/~amitp/yapps/
|
||||
# hasn't been touched since 2003 and has not been updated to work with
|
||||
# Python 3.
|
||||
#
|
||||
# 2.) Generate the grammar:
|
||||
#
|
||||
# yapps2 xpathparser.g xpathparser.py.proto
|
||||
#
|
||||
# 3.) Edit the output to depend on the embedded runtime, and remove extraneous
|
||||
# imports:
|
||||
#
|
||||
# sed -e '/^# Begin/,${/^[^ ].*mport/d}' -e 's/runtime\.//g' \
|
||||
# -e "s/^\(from __future\)/exec(r'''\n\1/" -e"\$a''')"
|
||||
# xpathparser.py.proto > xpathparser.py
|
||||
|
||||
"""
|
||||
XPath Parser.
|
||||
|
||||
Besides the parser code produced by Yapps, this module also defines the
|
||||
parse-time exception classes, a scanner class, a base class for parsers
|
||||
produced by Yapps, and a context class that keeps track of the parse stack.
|
||||
These have been copied from the Yapps runtime module.
|
||||
"""
|
||||
|
||||
from __future__ import print_function
|
||||
import sys, re
|
||||
|
||||
MIN_WINDOW=4096
|
||||
# File lookup window
|
||||
|
||||
class SyntaxError(Exception):
|
||||
"""When we run into an unexpected token, this is the exception to use"""
|
||||
def __init__(self, pos=None, msg="Bad Token", context=None):
|
||||
Exception.__init__(self)
|
||||
self.pos = pos
|
||||
self.msg = msg
|
||||
self.context = context
|
||||
|
||||
def __str__(self):
|
||||
if not self.pos: return 'SyntaxError'
|
||||
else: return 'SyntaxError@%s(%s)' % (repr(self.pos), self.msg)
|
||||
|
||||
class NoMoreTokens(Exception):
|
||||
"""Another exception object, for when we run out of tokens"""
|
||||
pass
|
||||
|
||||
class Token(object):
|
||||
"""Yapps token.
|
||||
|
||||
This is a container for a scanned token.
|
||||
"""
|
||||
|
||||
def __init__(self, type,value, pos=None):
|
||||
"""Initialize a token."""
|
||||
self.type = type
|
||||
self.value = value
|
||||
self.pos = pos
|
||||
|
||||
def __repr__(self):
|
||||
output = '<%s: %s' % (self.type, repr(self.value))
|
||||
if self.pos:
|
||||
output += " @ "
|
||||
if self.pos[0]:
|
||||
output += "%s:" % self.pos[0]
|
||||
if self.pos[1]:
|
||||
output += "%d" % self.pos[1]
|
||||
if self.pos[2] is not None:
|
||||
output += ".%d" % self.pos[2]
|
||||
output += ">"
|
||||
return output
|
||||
|
||||
in_name=0
|
||||
class Scanner(object):
|
||||
"""Yapps scanner.
|
||||
|
||||
The Yapps scanner can work in context sensitive or context
|
||||
insensitive modes. The token(i) method is used to retrieve the
|
||||
i-th token. It takes a restrict set that limits the set of tokens
|
||||
it is allowed to return. In context sensitive mode, this restrict
|
||||
set guides the scanner. In context insensitive mode, there is no
|
||||
restriction (the set is always the full set of tokens).
|
||||
|
||||
"""
|
||||
|
||||
def __init__(self, patterns, ignore, input="",
|
||||
file=None,filename=None,stacked=False):
|
||||
"""Initialize the scanner.
|
||||
|
||||
Parameters:
|
||||
patterns : [(terminal, uncompiled regex), ...] or None
|
||||
ignore : {terminal:None, ...}
|
||||
input : string
|
||||
|
||||
If patterns is None, we assume that the subclass has
|
||||
defined self.patterns : [(terminal, compiled regex), ...].
|
||||
Note that the patterns parameter expects uncompiled regexes,
|
||||
whereas the self.patterns field expects compiled regexes.
|
||||
|
||||
The 'ignore' value is either None or a callable, which is called
|
||||
with the scanner and the to-be-ignored match object; this can
|
||||
be used for include file or comment handling.
|
||||
"""
|
||||
|
||||
if not filename:
|
||||
global in_name
|
||||
filename="<f.%d>" % in_name
|
||||
in_name += 1
|
||||
|
||||
self.input = input
|
||||
self.ignore = ignore
|
||||
self.file = file
|
||||
self.filename = filename
|
||||
self.pos = 0
|
||||
self.del_pos = 0 # skipped
|
||||
self.line = 1
|
||||
self.del_line = 0 # skipped
|
||||
self.col = 0
|
||||
self.tokens = []
|
||||
self.stack = None
|
||||
self.stacked = stacked
|
||||
|
||||
self.last_read_token = None
|
||||
self.last_token = None
|
||||
self.last_types = None
|
||||
|
||||
if patterns is not None:
|
||||
# Compile the regex strings into regex objects
|
||||
self.patterns = []
|
||||
for terminal, regex in patterns:
|
||||
self.patterns.append( (terminal, re.compile(regex)) )
|
||||
|
||||
def stack_input(self, input="", file=None, filename=None):
|
||||
"""Temporarily parse from a second file."""
|
||||
|
||||
# Already reading from somewhere else: Go on top of that, please.
|
||||
if self.stack:
|
||||
# autogenerate a recursion-level-identifying filename
|
||||
if not filename:
|
||||
filename = 1
|
||||
else:
|
||||
try:
|
||||
filename += 1
|
||||
except TypeError:
|
||||
pass
|
||||
# now pass off to the include file
|
||||
self.stack.stack_input(input,file,filename)
|
||||
else:
|
||||
|
||||
try:
|
||||
filename += 0
|
||||
except TypeError:
|
||||
pass
|
||||
else:
|
||||
filename = "<str_%d>" % filename
|
||||
|
||||
# self.stack = object.__new__(self.__class__)
|
||||
# Scanner.__init__(self.stack,self.patterns,self.ignore,input,file,filename, stacked=True)
|
||||
|
||||
# Note that the pattern+ignore are added by the generated
|
||||
# scanner code
|
||||
self.stack = self.__class__(input,file,filename, stacked=True)
|
||||
|
||||
def get_pos(self):
|
||||
"""Return a file/line/char tuple."""
|
||||
if self.stack: return self.stack.get_pos()
|
||||
|
||||
return (self.filename, self.line+self.del_line, self.col)
|
||||
|
||||
# def __repr__(self):
|
||||
# """Print the last few tokens that have been scanned in"""
|
||||
# output = ''
|
||||
# for t in self.tokens:
|
||||
# output += '%s\n' % (repr(t),)
|
||||
# return output
|
||||
|
||||
def print_line_with_pointer(self, pos, length=0, out=sys.stderr):
|
||||
"""Print the line of 'text' that includes position 'p',
|
||||
along with a second line with a single caret (^) at position p"""
|
||||
|
||||
file,line,p = pos
|
||||
if file != self.filename:
|
||||
if self.stack: return self.stack.print_line_with_pointer(pos,length=length,out=out)
|
||||
print >>out, "(%s: not in input buffer)" % file
|
||||
return
|
||||
|
||||
text = self.input
|
||||
p += length-1 # starts at pos 1
|
||||
|
||||
origline=line
|
||||
line -= self.del_line
|
||||
spos=0
|
||||
if line > 0:
|
||||
while 1:
|
||||
line = line - 1
|
||||
try:
|
||||
cr = text.index("\n",spos)
|
||||
except ValueError:
|
||||
if line:
|
||||
text = ""
|
||||
break
|
||||
if line == 0:
|
||||
text = text[spos:cr]
|
||||
break
|
||||
spos = cr+1
|
||||
else:
|
||||
print >>out, "(%s:%d not in input buffer)" % (file,origline)
|
||||
return
|
||||
|
||||
# Now try printing part of the line
|
||||
text = text[max(p-80, 0):p+80]
|
||||
p = p - max(p-80, 0)
|
||||
|
||||
# Strip to the left
|
||||
i = text[:p].rfind('\n')
|
||||
j = text[:p].rfind('\r')
|
||||
if i < 0 or (0 <= j < i): i = j
|
||||
if 0 <= i < p:
|
||||
p = p - i - 1
|
||||
text = text[i+1:]
|
||||
|
||||
# Strip to the right
|
||||
i = text.find('\n', p)
|
||||
j = text.find('\r', p)
|
||||
if i < 0 or (0 <= j < i): i = j
|
||||
if i >= 0:
|
||||
text = text[:i]
|
||||
|
||||
# Now shorten the text
|
||||
while len(text) > 70 and p > 60:
|
||||
# Cut off 10 chars
|
||||
text = "..." + text[10:]
|
||||
p = p - 7
|
||||
|
||||
# Now print the string, along with an indicator
|
||||
print >>out, '> ',text
|
||||
print >>out, '> ',' '*p + '^'
|
||||
|
||||
def grab_input(self):
|
||||
"""Get more input if possible."""
|
||||
if not self.file: return
|
||||
if len(self.input) - self.pos >= MIN_WINDOW: return
|
||||
|
||||
data = self.file.read(MIN_WINDOW)
|
||||
if data is None or data == "":
|
||||
self.file = None
|
||||
|
||||
# Drop bytes from the start, if necessary.
|
||||
if self.pos > 2*MIN_WINDOW:
|
||||
self.del_pos += MIN_WINDOW
|
||||
self.del_line += self.input[:MIN_WINDOW].count("\n")
|
||||
self.pos -= MIN_WINDOW
|
||||
self.input = self.input[MIN_WINDOW:] + data
|
||||
else:
|
||||
self.input = self.input + data
|
||||
|
||||
def getchar(self):
|
||||
"""Return the next character."""
|
||||
self.grab_input()
|
||||
|
||||
c = self.input[self.pos]
|
||||
self.pos += 1
|
||||
return c
|
||||
|
||||
def token(self, restrict, context=None):
|
||||
"""Scan for another token."""
|
||||
|
||||
while 1:
|
||||
if self.stack:
|
||||
try:
|
||||
return self.stack.token(restrict, context)
|
||||
except StopIteration:
|
||||
self.stack = None
|
||||
|
||||
# Keep looking for a token, ignoring any in self.ignore
|
||||
self.grab_input()
|
||||
|
||||
# special handling for end-of-file
|
||||
if self.stacked and self.pos==len(self.input):
|
||||
raise StopIteration
|
||||
|
||||
# Search the patterns for the longest match, with earlier
|
||||
# tokens in the list having preference
|
||||
best_match = -1
|
||||
best_pat = '(error)'
|
||||
best_m = None
|
||||
for p, regexp in self.patterns:
|
||||
# First check to see if we're ignoring this token
|
||||
if restrict and p not in restrict and p not in self.ignore:
|
||||
continue
|
||||
m = regexp.match(self.input, self.pos)
|
||||
if m and m.end()-m.start() > best_match:
|
||||
# We got a match that's better than the previous one
|
||||
best_pat = p
|
||||
best_match = m.end()-m.start()
|
||||
best_m = m
|
||||
|
||||
# If we didn't find anything, raise an error
|
||||
if best_pat == '(error)' and best_match < 0:
|
||||
msg = 'Bad Token'
|
||||
if restrict:
|
||||
msg = 'Trying to find one of '+', '.join(restrict)
|
||||
raise SyntaxError(self.get_pos(), msg, context=context)
|
||||
|
||||
ignore = best_pat in self.ignore
|
||||
value = self.input[self.pos:self.pos+best_match]
|
||||
if not ignore:
|
||||
tok=Token(type=best_pat, value=value, pos=self.get_pos())
|
||||
|
||||
self.pos += best_match
|
||||
|
||||
npos = value.rfind("\n")
|
||||
if npos > -1:
|
||||
self.col = best_match-npos
|
||||
self.line += value.count("\n")
|
||||
else:
|
||||
self.col += best_match
|
||||
|
||||
# If we found something that isn't to be ignored, return it
|
||||
if not ignore:
|
||||
if len(self.tokens) >= 10:
|
||||
del self.tokens[0]
|
||||
self.tokens.append(tok)
|
||||
self.last_read_token = tok
|
||||
# print repr(tok)
|
||||
return tok
|
||||
else:
|
||||
ignore = self.ignore[best_pat]
|
||||
if ignore:
|
||||
ignore(self, best_m)
|
||||
|
||||
def peek(self, *types, **kw):
|
||||
"""Returns the token type for lookahead; if there are any args
|
||||
then the list of args is the set of token types to allow"""
|
||||
context = kw.get("context",None)
|
||||
if self.last_token is None:
|
||||
self.last_types = types
|
||||
self.last_token = self.token(types,context)
|
||||
elif self.last_types:
|
||||
for t in types:
|
||||
if t not in self.last_types:
|
||||
raise NotImplementedError("Unimplemented: restriction set changed")
|
||||
return self.last_token.type
|
||||
|
||||
def scan(self, type, **kw):
|
||||
"""Returns the matched text, and moves to the next token"""
|
||||
context = kw.get("context",None)
|
||||
|
||||
if self.last_token is None:
|
||||
tok = self.token([type],context)
|
||||
else:
|
||||
if self.last_types and type not in self.last_types:
|
||||
raise NotImplementedError("Unimplemented: restriction set changed")
|
||||
|
||||
tok = self.last_token
|
||||
self.last_token = None
|
||||
if tok.type != type:
|
||||
if not self.last_types: self.last_types=[]
|
||||
raise SyntaxError(tok.pos, 'Trying to find '+type+': '+ ', '.join(self.last_types)+", got "+tok.type, context=context)
|
||||
return tok.value
|
||||
|
||||
class Parser(object):
|
||||
"""Base class for Yapps-generated parsers.
|
||||
|
||||
"""
|
||||
|
||||
def __init__(self, scanner):
|
||||
self._scanner = scanner
|
||||
|
||||
def _stack(self, input="",file=None,filename=None):
|
||||
"""Temporarily read from someplace else"""
|
||||
self._scanner.stack_input(input,file,filename)
|
||||
self._tok = None
|
||||
|
||||
def _peek(self, *types, **kw):
|
||||
"""Returns the token type for lookahead; if there are any args
|
||||
then the list of args is the set of token types to allow"""
|
||||
return self._scanner.peek(*types, **kw)
|
||||
|
||||
def _scan(self, type, **kw):
|
||||
"""Returns the matched text, and moves to the next token"""
|
||||
return self._scanner.scan(type, **kw)
|
||||
|
||||
class Context(object):
|
||||
"""Class to represent the parser's call stack.
|
||||
|
||||
Every rule creates a Context that links to its parent rule. The
|
||||
contexts can be used for debugging.
|
||||
|
||||
"""
|
||||
|
||||
def __init__(self, parent, scanner, rule, args=()):
|
||||
"""Create a new context.
|
||||
|
||||
Args:
|
||||
parent: Context object or None
|
||||
scanner: Scanner object
|
||||
rule: string (name of the rule)
|
||||
args: tuple listing parameters to the rule
|
||||
|
||||
"""
|
||||
self.parent = parent
|
||||
self.scanner = scanner
|
||||
self.rule = rule
|
||||
self.args = args
|
||||
while scanner.stack: scanner = scanner.stack
|
||||
self.token = scanner.last_read_token
|
||||
|
||||
def __str__(self):
|
||||
output = ''
|
||||
if self.parent: output = str(self.parent) + ' > '
|
||||
output += self.rule
|
||||
return output
|
||||
|
||||
def print_error(err, scanner, max_ctx=None):
|
||||
"""Print error messages, the parser stack, and the input text -- for human-readable error messages."""
|
||||
# NOTE: this function assumes 80 columns :-(
|
||||
# Figure out the line number
|
||||
pos = err.pos
|
||||
if not pos:
|
||||
pos = scanner.get_pos()
|
||||
|
||||
file_name, line_number, column_number = pos
|
||||
print('%s:%d:%d: %s' % (file_name, line_number, column_number, err.msg), file=sys.stderr)
|
||||
|
||||
scanner.print_line_with_pointer(pos)
|
||||
|
||||
context = err.context
|
||||
token = None
|
||||
while context:
|
||||
print('while parsing %s%s:' % (context.rule, tuple(context.args)), file=sys.stderr)
|
||||
if context.token:
|
||||
token = context.token
|
||||
if token:
|
||||
scanner.print_line_with_pointer(token.pos, length=len(token.value))
|
||||
context = context.parent
|
||||
if max_ctx:
|
||||
max_ctx = max_ctx-1
|
||||
if not max_ctx:
|
||||
break
|
||||
|
||||
def wrap_error_reporter(parser, rule, *args,**kw):
|
||||
try:
|
||||
return getattr(parser, rule)(*args,**kw)
|
||||
except SyntaxError as e:
|
||||
print_error(e, parser._scanner)
|
||||
except NoMoreTokens:
|
||||
print('Could not complete parsing; stopped around here:', file=sys.stderr)
|
||||
print(parser._scanner, file=sys.stderr)
|
||||
|
||||
from twisted.words.xish.xpath import AttribValue, BooleanValue, CompareValue
|
||||
from twisted.words.xish.xpath import Function, IndexValue, LiteralValue
|
||||
from twisted.words.xish.xpath import _AnyLocation, _Location
|
||||
|
||||
%%
|
||||
parser XPathParser:
|
||||
ignore: "\\s+"
|
||||
token INDEX: "[0-9]+"
|
||||
token WILDCARD: "\*"
|
||||
token IDENTIFIER: "[a-zA-Z][a-zA-Z0-9_\-]*"
|
||||
token ATTRIBUTE: "\@[a-zA-Z][a-zA-Z0-9_\-]*"
|
||||
token FUNCNAME: "[a-zA-Z][a-zA-Z0-9_]*"
|
||||
token CMP_EQ: "\="
|
||||
token CMP_NE: "\!\="
|
||||
token STR_DQ: '"([^"]|(\\"))*?"'
|
||||
token STR_SQ: "'([^']|(\\'))*?'"
|
||||
token OP_AND: "and"
|
||||
token OP_OR: "or"
|
||||
token END: "$"
|
||||
|
||||
rule XPATH: PATH {{ result = PATH; current = result }}
|
||||
( PATH {{ current.childLocation = PATH; current = current.childLocation }} ) * END
|
||||
{{ return result }}
|
||||
|
||||
rule PATH: ("/" {{ result = _Location() }} | "//" {{ result = _AnyLocation() }} )
|
||||
( IDENTIFIER {{ result.elementName = IDENTIFIER }} | WILDCARD {{ result.elementName = None }} )
|
||||
( "\[" PREDICATE {{ result.predicates.append(PREDICATE) }} "\]")*
|
||||
{{ return result }}
|
||||
|
||||
rule PREDICATE: EXPR {{ return EXPR }} |
|
||||
INDEX {{ return IndexValue(INDEX) }}
|
||||
|
||||
rule EXPR: FACTOR {{ e = FACTOR }}
|
||||
( BOOLOP FACTOR {{ e = BooleanValue(e, BOOLOP, FACTOR) }} )*
|
||||
{{ return e }}
|
||||
|
||||
rule BOOLOP: ( OP_AND {{ return OP_AND }} | OP_OR {{ return OP_OR }} )
|
||||
|
||||
rule FACTOR: TERM {{ return TERM }}
|
||||
| "\(" EXPR "\)" {{ return EXPR }}
|
||||
|
||||
rule TERM: VALUE {{ t = VALUE }}
|
||||
[ CMP VALUE {{ t = CompareValue(t, CMP, VALUE) }} ]
|
||||
{{ return t }}
|
||||
|
||||
rule VALUE: "@" IDENTIFIER {{ return AttribValue(IDENTIFIER) }} |
|
||||
FUNCNAME {{ f = Function(FUNCNAME); args = [] }}
|
||||
"\(" [ VALUE {{ args.append(VALUE) }}
|
||||
(
|
||||
"," VALUE {{ args.append(VALUE) }}
|
||||
)*
|
||||
] "\)" {{ f.setParams(*args); return f }} |
|
||||
STR {{ return LiteralValue(STR[1:len(STR)-1]) }}
|
||||
|
||||
rule CMP: (CMP_EQ {{ return CMP_EQ }} | CMP_NE {{ return CMP_NE }})
|
||||
rule STR: (STR_DQ {{ return STR_DQ }} | STR_SQ {{ return STR_SQ }})
|
||||
@@ -0,0 +1,650 @@
|
||||
# -*- test-case-name: twisted.words.test.test_xpath -*-
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
# pylint: disable=W9401,W9402
|
||||
|
||||
# DO NOT EDIT xpathparser.py!
|
||||
#
|
||||
# It is generated from xpathparser.g using Yapps. Make needed changes there.
|
||||
# This also means that the generated Python may not conform to Twisted's coding
|
||||
# standards, so it is wrapped in exec to prevent automated checkers from
|
||||
# complaining.
|
||||
|
||||
# HOWTO Generate me:
|
||||
#
|
||||
# 1.) Grab a copy of yapps2:
|
||||
# https://github.com/smurfix/yapps
|
||||
#
|
||||
# Note: Do NOT use the package in debian/ubuntu as it has incompatible
|
||||
# modifications. The original at http://theory.stanford.edu/~amitp/yapps/
|
||||
# hasn't been touched since 2003 and has not been updated to work with
|
||||
# Python 3.
|
||||
#
|
||||
# 2.) Generate the grammar:
|
||||
#
|
||||
# yapps2 xpathparser.g xpathparser.py.proto
|
||||
#
|
||||
# 3.) Edit the output to depend on the embedded runtime, and remove extraneous
|
||||
# imports:
|
||||
#
|
||||
# sed -e '/^# Begin/,${/^[^ ].*mport/d}' -e '/^[^#]/s/runtime\.//g' \
|
||||
# -e "s/^\(from __future\)/exec(r'''\n\1/" -e"\$a''')"
|
||||
# xpathparser.py.proto > xpathparser.py
|
||||
|
||||
"""
|
||||
XPath Parser.
|
||||
|
||||
Besides the parser code produced by Yapps, this module also defines the
|
||||
parse-time exception classes, a scanner class, a base class for parsers
|
||||
produced by Yapps, and a context class that keeps track of the parse stack.
|
||||
These have been copied from the Yapps runtime module.
|
||||
"""
|
||||
|
||||
exec(r'''
|
||||
from __future__ import print_function
|
||||
import sys, re
|
||||
|
||||
MIN_WINDOW=4096
|
||||
# File lookup window
|
||||
|
||||
class SyntaxError(Exception):
|
||||
"""When we run into an unexpected token, this is the exception to use"""
|
||||
def __init__(self, pos=None, msg="Bad Token", context=None):
|
||||
Exception.__init__(self)
|
||||
self.pos = pos
|
||||
self.msg = msg
|
||||
self.context = context
|
||||
|
||||
def __str__(self):
|
||||
if not self.pos: return 'SyntaxError'
|
||||
else: return 'SyntaxError@%s(%s)' % (repr(self.pos), self.msg)
|
||||
|
||||
class NoMoreTokens(Exception):
|
||||
"""Another exception object, for when we run out of tokens"""
|
||||
pass
|
||||
|
||||
class Token(object):
|
||||
"""Yapps token.
|
||||
|
||||
This is a container for a scanned token.
|
||||
"""
|
||||
|
||||
def __init__(self, type,value, pos=None):
|
||||
"""Initialize a token."""
|
||||
self.type = type
|
||||
self.value = value
|
||||
self.pos = pos
|
||||
|
||||
def __repr__(self):
|
||||
output = '<%s: %s' % (self.type, repr(self.value))
|
||||
if self.pos:
|
||||
output += " @ "
|
||||
if self.pos[0]:
|
||||
output += "%s:" % self.pos[0]
|
||||
if self.pos[1]:
|
||||
output += "%d" % self.pos[1]
|
||||
if self.pos[2] is not None:
|
||||
output += ".%d" % self.pos[2]
|
||||
output += ">"
|
||||
return output
|
||||
|
||||
in_name=0
|
||||
class Scanner(object):
|
||||
"""Yapps scanner.
|
||||
|
||||
The Yapps scanner can work in context sensitive or context
|
||||
insensitive modes. The token(i) method is used to retrieve the
|
||||
i-th token. It takes a restrict set that limits the set of tokens
|
||||
it is allowed to return. In context sensitive mode, this restrict
|
||||
set guides the scanner. In context insensitive mode, there is no
|
||||
restriction (the set is always the full set of tokens).
|
||||
|
||||
"""
|
||||
|
||||
def __init__(self, patterns, ignore, input="",
|
||||
file=None,filename=None,stacked=False):
|
||||
"""Initialize the scanner.
|
||||
|
||||
Parameters:
|
||||
patterns : [(terminal, uncompiled regex), ...] or None
|
||||
ignore : {terminal:None, ...}
|
||||
input : string
|
||||
|
||||
If patterns is None, we assume that the subclass has
|
||||
defined self.patterns : [(terminal, compiled regex), ...].
|
||||
Note that the patterns parameter expects uncompiled regexes,
|
||||
whereas the self.patterns field expects compiled regexes.
|
||||
|
||||
The 'ignore' value is either None or a callable, which is called
|
||||
with the scanner and the to-be-ignored match object; this can
|
||||
be used for include file or comment handling.
|
||||
"""
|
||||
|
||||
if not filename:
|
||||
global in_name
|
||||
filename="<f.%d>" % in_name
|
||||
in_name += 1
|
||||
|
||||
self.input = input
|
||||
self.ignore = ignore
|
||||
self.file = file
|
||||
self.filename = filename
|
||||
self.pos = 0
|
||||
self.del_pos = 0 # skipped
|
||||
self.line = 1
|
||||
self.del_line = 0 # skipped
|
||||
self.col = 0
|
||||
self.tokens = []
|
||||
self.stack = None
|
||||
self.stacked = stacked
|
||||
|
||||
self.last_read_token = None
|
||||
self.last_token = None
|
||||
self.last_types = None
|
||||
|
||||
if patterns is not None:
|
||||
# Compile the regex strings into regex objects
|
||||
self.patterns = []
|
||||
for terminal, regex in patterns:
|
||||
self.patterns.append( (terminal, re.compile(regex)) )
|
||||
|
||||
def stack_input(self, input="", file=None, filename=None):
|
||||
"""Temporarily parse from a second file."""
|
||||
|
||||
# Already reading from somewhere else: Go on top of that, please.
|
||||
if self.stack:
|
||||
# autogenerate a recursion-level-identifying filename
|
||||
if not filename:
|
||||
filename = 1
|
||||
else:
|
||||
try:
|
||||
filename += 1
|
||||
except TypeError:
|
||||
pass
|
||||
# now pass off to the include file
|
||||
self.stack.stack_input(input,file,filename)
|
||||
else:
|
||||
|
||||
try:
|
||||
filename += 0
|
||||
except TypeError:
|
||||
pass
|
||||
else:
|
||||
filename = "<str_%d>" % filename
|
||||
|
||||
# self.stack = object.__new__(self.__class__)
|
||||
# Scanner.__init__(self.stack,self.patterns,self.ignore,input,file,filename, stacked=True)
|
||||
|
||||
# Note that the pattern+ignore are added by the generated
|
||||
# scanner code
|
||||
self.stack = self.__class__(input,file,filename, stacked=True)
|
||||
|
||||
def get_pos(self):
|
||||
"""Return a file/line/char tuple."""
|
||||
if self.stack: return self.stack.get_pos()
|
||||
|
||||
return (self.filename, self.line+self.del_line, self.col)
|
||||
|
||||
# def __repr__(self):
|
||||
# """Print the last few tokens that have been scanned in"""
|
||||
# output = ''
|
||||
# for t in self.tokens:
|
||||
# output += '%s\n' % (repr(t),)
|
||||
# return output
|
||||
|
||||
def print_line_with_pointer(self, pos, length=0, out=sys.stderr):
|
||||
"""Print the line of 'text' that includes position 'p',
|
||||
along with a second line with a single caret (^) at position p"""
|
||||
|
||||
file,line,p = pos
|
||||
if file != self.filename:
|
||||
if self.stack: return self.stack.print_line_with_pointer(pos,length=length,out=out)
|
||||
print >>out, "(%s: not in input buffer)" % file
|
||||
return
|
||||
|
||||
text = self.input
|
||||
p += length-1 # starts at pos 1
|
||||
|
||||
origline=line
|
||||
line -= self.del_line
|
||||
spos=0
|
||||
if line > 0:
|
||||
while 1:
|
||||
line = line - 1
|
||||
try:
|
||||
cr = text.index("\n",spos)
|
||||
except ValueError:
|
||||
if line:
|
||||
text = ""
|
||||
break
|
||||
if line == 0:
|
||||
text = text[spos:cr]
|
||||
break
|
||||
spos = cr+1
|
||||
else:
|
||||
print >>out, "(%s:%d not in input buffer)" % (file,origline)
|
||||
return
|
||||
|
||||
# Now try printing part of the line
|
||||
text = text[max(p-80, 0):p+80]
|
||||
p = p - max(p-80, 0)
|
||||
|
||||
# Strip to the left
|
||||
i = text[:p].rfind('\n')
|
||||
j = text[:p].rfind('\r')
|
||||
if i < 0 or (0 <= j < i): i = j
|
||||
if 0 <= i < p:
|
||||
p = p - i - 1
|
||||
text = text[i+1:]
|
||||
|
||||
# Strip to the right
|
||||
i = text.find('\n', p)
|
||||
j = text.find('\r', p)
|
||||
if i < 0 or (0 <= j < i): i = j
|
||||
if i >= 0:
|
||||
text = text[:i]
|
||||
|
||||
# Now shorten the text
|
||||
while len(text) > 70 and p > 60:
|
||||
# Cut off 10 chars
|
||||
text = "..." + text[10:]
|
||||
p = p - 7
|
||||
|
||||
# Now print the string, along with an indicator
|
||||
print >>out, '> ',text
|
||||
print >>out, '> ',' '*p + '^'
|
||||
|
||||
def grab_input(self):
|
||||
"""Get more input if possible."""
|
||||
if not self.file: return
|
||||
if len(self.input) - self.pos >= MIN_WINDOW: return
|
||||
|
||||
data = self.file.read(MIN_WINDOW)
|
||||
if data is None or data == "":
|
||||
self.file = None
|
||||
|
||||
# Drop bytes from the start, if necessary.
|
||||
if self.pos > 2*MIN_WINDOW:
|
||||
self.del_pos += MIN_WINDOW
|
||||
self.del_line += self.input[:MIN_WINDOW].count("\n")
|
||||
self.pos -= MIN_WINDOW
|
||||
self.input = self.input[MIN_WINDOW:] + data
|
||||
else:
|
||||
self.input = self.input + data
|
||||
|
||||
def getchar(self):
|
||||
"""Return the next character."""
|
||||
self.grab_input()
|
||||
|
||||
c = self.input[self.pos]
|
||||
self.pos += 1
|
||||
return c
|
||||
|
||||
def token(self, restrict, context=None):
|
||||
"""Scan for another token."""
|
||||
|
||||
while 1:
|
||||
if self.stack:
|
||||
try:
|
||||
return self.stack.token(restrict, context)
|
||||
except StopIteration:
|
||||
self.stack = None
|
||||
|
||||
# Keep looking for a token, ignoring any in self.ignore
|
||||
self.grab_input()
|
||||
|
||||
# special handling for end-of-file
|
||||
if self.stacked and self.pos==len(self.input):
|
||||
raise StopIteration
|
||||
|
||||
# Search the patterns for the longest match, with earlier
|
||||
# tokens in the list having preference
|
||||
best_match = -1
|
||||
best_pat = '(error)'
|
||||
best_m = None
|
||||
for p, regexp in self.patterns:
|
||||
# First check to see if we're ignoring this token
|
||||
if restrict and p not in restrict and p not in self.ignore:
|
||||
continue
|
||||
m = regexp.match(self.input, self.pos)
|
||||
if m and m.end()-m.start() > best_match:
|
||||
# We got a match that's better than the previous one
|
||||
best_pat = p
|
||||
best_match = m.end()-m.start()
|
||||
best_m = m
|
||||
|
||||
# If we didn't find anything, raise an error
|
||||
if best_pat == '(error)' and best_match < 0:
|
||||
msg = 'Bad Token'
|
||||
if restrict:
|
||||
msg = 'Trying to find one of '+', '.join(restrict)
|
||||
raise SyntaxError(self.get_pos(), msg, context=context)
|
||||
|
||||
ignore = best_pat in self.ignore
|
||||
value = self.input[self.pos:self.pos+best_match]
|
||||
if not ignore:
|
||||
tok=Token(type=best_pat, value=value, pos=self.get_pos())
|
||||
|
||||
self.pos += best_match
|
||||
|
||||
npos = value.rfind("\n")
|
||||
if npos > -1:
|
||||
self.col = best_match-npos
|
||||
self.line += value.count("\n")
|
||||
else:
|
||||
self.col += best_match
|
||||
|
||||
# If we found something that isn't to be ignored, return it
|
||||
if not ignore:
|
||||
if len(self.tokens) >= 10:
|
||||
del self.tokens[0]
|
||||
self.tokens.append(tok)
|
||||
self.last_read_token = tok
|
||||
# print repr(tok)
|
||||
return tok
|
||||
else:
|
||||
ignore = self.ignore[best_pat]
|
||||
if ignore:
|
||||
ignore(self, best_m)
|
||||
|
||||
def peek(self, *types, **kw):
|
||||
"""Returns the token type for lookahead; if there are any args
|
||||
then the list of args is the set of token types to allow"""
|
||||
context = kw.get("context",None)
|
||||
if self.last_token is None:
|
||||
self.last_types = types
|
||||
self.last_token = self.token(types,context)
|
||||
elif self.last_types:
|
||||
for t in types:
|
||||
if t not in self.last_types:
|
||||
raise NotImplementedError("Unimplemented: restriction set changed")
|
||||
return self.last_token.type
|
||||
|
||||
def scan(self, type, **kw):
|
||||
"""Returns the matched text, and moves to the next token"""
|
||||
context = kw.get("context",None)
|
||||
|
||||
if self.last_token is None:
|
||||
tok = self.token([type],context)
|
||||
else:
|
||||
if self.last_types and type not in self.last_types:
|
||||
raise NotImplementedError("Unimplemented: restriction set changed")
|
||||
|
||||
tok = self.last_token
|
||||
self.last_token = None
|
||||
if tok.type != type:
|
||||
if not self.last_types: self.last_types=[]
|
||||
raise SyntaxError(tok.pos, 'Trying to find '+type+': '+ ', '.join(self.last_types)+", got "+tok.type, context=context)
|
||||
return tok.value
|
||||
|
||||
class Parser(object):
|
||||
"""Base class for Yapps-generated parsers.
|
||||
|
||||
"""
|
||||
|
||||
def __init__(self, scanner):
|
||||
self._scanner = scanner
|
||||
|
||||
def _stack(self, input="",file=None,filename=None):
|
||||
"""Temporarily read from someplace else"""
|
||||
self._scanner.stack_input(input,file,filename)
|
||||
self._tok = None
|
||||
|
||||
def _peek(self, *types, **kw):
|
||||
"""Returns the token type for lookahead; if there are any args
|
||||
then the list of args is the set of token types to allow"""
|
||||
return self._scanner.peek(*types, **kw)
|
||||
|
||||
def _scan(self, type, **kw):
|
||||
"""Returns the matched text, and moves to the next token"""
|
||||
return self._scanner.scan(type, **kw)
|
||||
|
||||
class Context(object):
|
||||
"""Class to represent the parser's call stack.
|
||||
|
||||
Every rule creates a Context that links to its parent rule. The
|
||||
contexts can be used for debugging.
|
||||
|
||||
"""
|
||||
|
||||
def __init__(self, parent, scanner, rule, args=()):
|
||||
"""Create a new context.
|
||||
|
||||
Args:
|
||||
parent: Context object or None
|
||||
scanner: Scanner object
|
||||
rule: string (name of the rule)
|
||||
args: tuple listing parameters to the rule
|
||||
|
||||
"""
|
||||
self.parent = parent
|
||||
self.scanner = scanner
|
||||
self.rule = rule
|
||||
self.args = args
|
||||
while scanner.stack: scanner = scanner.stack
|
||||
self.token = scanner.last_read_token
|
||||
|
||||
def __str__(self):
|
||||
output = ''
|
||||
if self.parent: output = str(self.parent) + ' > '
|
||||
output += self.rule
|
||||
return output
|
||||
|
||||
def print_error(err, scanner, max_ctx=None):
|
||||
"""Print error messages, the parser stack, and the input text -- for human-readable error messages."""
|
||||
# NOTE: this function assumes 80 columns :-(
|
||||
# Figure out the line number
|
||||
pos = err.pos
|
||||
if not pos:
|
||||
pos = scanner.get_pos()
|
||||
|
||||
file_name, line_number, column_number = pos
|
||||
print('%s:%d:%d: %s' % (file_name, line_number, column_number, err.msg), file=sys.stderr)
|
||||
|
||||
scanner.print_line_with_pointer(pos)
|
||||
|
||||
context = err.context
|
||||
token = None
|
||||
while context:
|
||||
print('while parsing %s%s:' % (context.rule, tuple(context.args)), file=sys.stderr)
|
||||
if context.token:
|
||||
token = context.token
|
||||
if token:
|
||||
scanner.print_line_with_pointer(token.pos, length=len(token.value))
|
||||
context = context.parent
|
||||
if max_ctx:
|
||||
max_ctx = max_ctx-1
|
||||
if not max_ctx:
|
||||
break
|
||||
|
||||
def wrap_error_reporter(parser, rule, *args,**kw):
|
||||
try:
|
||||
return getattr(parser, rule)(*args,**kw)
|
||||
except SyntaxError as e:
|
||||
print_error(e, parser._scanner)
|
||||
except NoMoreTokens:
|
||||
print('Could not complete parsing; stopped around here:', file=sys.stderr)
|
||||
print(parser._scanner, file=sys.stderr)
|
||||
|
||||
from twisted.words.xish.xpath import AttribValue, BooleanValue, CompareValue
|
||||
from twisted.words.xish.xpath import Function, IndexValue, LiteralValue
|
||||
from twisted.words.xish.xpath import _AnyLocation, _Location
|
||||
|
||||
|
||||
# Begin -- grammar generated by Yapps
|
||||
|
||||
class XPathParserScanner(Scanner):
|
||||
patterns = [
|
||||
('","', re.compile(',')),
|
||||
('"@"', re.compile('@')),
|
||||
('"\\)"', re.compile('\\)')),
|
||||
('"\\("', re.compile('\\(')),
|
||||
('"\\]"', re.compile('\\]')),
|
||||
('"\\["', re.compile('\\[')),
|
||||
('"//"', re.compile('//')),
|
||||
('"/"', re.compile('/')),
|
||||
('\\s+', re.compile('\\s+')),
|
||||
('INDEX', re.compile('[0-9]+')),
|
||||
('WILDCARD', re.compile('\\*')),
|
||||
('IDENTIFIER', re.compile('[a-zA-Z][a-zA-Z0-9_\\-]*')),
|
||||
('ATTRIBUTE', re.compile('\\@[a-zA-Z][a-zA-Z0-9_\\-]*')),
|
||||
('FUNCNAME', re.compile('[a-zA-Z][a-zA-Z0-9_]*')),
|
||||
('CMP_EQ', re.compile('\\=')),
|
||||
('CMP_NE', re.compile('\\!\\=')),
|
||||
('STR_DQ', re.compile('"([^"]|(\\"))*?"')),
|
||||
('STR_SQ', re.compile("'([^']|(\\'))*?'")),
|
||||
('OP_AND', re.compile('and')),
|
||||
('OP_OR', re.compile('or')),
|
||||
('END', re.compile('$')),
|
||||
]
|
||||
def __init__(self, str,*args,**kw):
|
||||
Scanner.__init__(self,None,{'\\s+':None,},str,*args,**kw)
|
||||
|
||||
class XPathParser(Parser):
|
||||
Context = Context
|
||||
def XPATH(self, _parent=None):
|
||||
_context = self.Context(_parent, self._scanner, 'XPATH', [])
|
||||
PATH = self.PATH(_context)
|
||||
result = PATH; current = result
|
||||
while self._peek('END', '"/"', '"//"', context=_context) != 'END':
|
||||
PATH = self.PATH(_context)
|
||||
current.childLocation = PATH; current = current.childLocation
|
||||
END = self._scan('END', context=_context)
|
||||
return result
|
||||
|
||||
def PATH(self, _parent=None):
|
||||
_context = self.Context(_parent, self._scanner, 'PATH', [])
|
||||
_token = self._peek('"/"', '"//"', context=_context)
|
||||
if _token == '"/"':
|
||||
self._scan('"/"', context=_context)
|
||||
result = _Location()
|
||||
else: # == '"//"'
|
||||
self._scan('"//"', context=_context)
|
||||
result = _AnyLocation()
|
||||
_token = self._peek('IDENTIFIER', 'WILDCARD', context=_context)
|
||||
if _token == 'IDENTIFIER':
|
||||
IDENTIFIER = self._scan('IDENTIFIER', context=_context)
|
||||
result.elementName = IDENTIFIER
|
||||
else: # == 'WILDCARD'
|
||||
WILDCARD = self._scan('WILDCARD', context=_context)
|
||||
result.elementName = None
|
||||
while self._peek('"\\["', 'END', '"/"', '"//"', context=_context) == '"\\["':
|
||||
self._scan('"\\["', context=_context)
|
||||
PREDICATE = self.PREDICATE(_context)
|
||||
result.predicates.append(PREDICATE)
|
||||
self._scan('"\\]"', context=_context)
|
||||
return result
|
||||
|
||||
def PREDICATE(self, _parent=None):
|
||||
_context = self.Context(_parent, self._scanner, 'PREDICATE', [])
|
||||
_token = self._peek('INDEX', '"\\("', '"@"', 'FUNCNAME', 'STR_DQ', 'STR_SQ', context=_context)
|
||||
if _token != 'INDEX':
|
||||
EXPR = self.EXPR(_context)
|
||||
return EXPR
|
||||
else: # == 'INDEX'
|
||||
INDEX = self._scan('INDEX', context=_context)
|
||||
return IndexValue(INDEX)
|
||||
|
||||
def EXPR(self, _parent=None):
|
||||
_context = self.Context(_parent, self._scanner, 'EXPR', [])
|
||||
FACTOR = self.FACTOR(_context)
|
||||
e = FACTOR
|
||||
while self._peek('OP_AND', 'OP_OR', '"\\)"', '"\\]"', context=_context) in ['OP_AND', 'OP_OR']:
|
||||
BOOLOP = self.BOOLOP(_context)
|
||||
FACTOR = self.FACTOR(_context)
|
||||
e = BooleanValue(e, BOOLOP, FACTOR)
|
||||
return e
|
||||
|
||||
def BOOLOP(self, _parent=None):
|
||||
_context = self.Context(_parent, self._scanner, 'BOOLOP', [])
|
||||
_token = self._peek('OP_AND', 'OP_OR', context=_context)
|
||||
if _token == 'OP_AND':
|
||||
OP_AND = self._scan('OP_AND', context=_context)
|
||||
return OP_AND
|
||||
else: # == 'OP_OR'
|
||||
OP_OR = self._scan('OP_OR', context=_context)
|
||||
return OP_OR
|
||||
|
||||
def FACTOR(self, _parent=None):
|
||||
_context = self.Context(_parent, self._scanner, 'FACTOR', [])
|
||||
_token = self._peek('"\\("', '"@"', 'FUNCNAME', 'STR_DQ', 'STR_SQ', context=_context)
|
||||
if _token != '"\\("':
|
||||
TERM = self.TERM(_context)
|
||||
return TERM
|
||||
else: # == '"\\("'
|
||||
self._scan('"\\("', context=_context)
|
||||
EXPR = self.EXPR(_context)
|
||||
self._scan('"\\)"', context=_context)
|
||||
return EXPR
|
||||
|
||||
def TERM(self, _parent=None):
|
||||
_context = self.Context(_parent, self._scanner, 'TERM', [])
|
||||
VALUE = self.VALUE(_context)
|
||||
t = VALUE
|
||||
if self._peek('CMP_EQ', 'CMP_NE', 'OP_AND', 'OP_OR', '"\\)"', '"\\]"', context=_context) in ['CMP_EQ', 'CMP_NE']:
|
||||
CMP = self.CMP(_context)
|
||||
VALUE = self.VALUE(_context)
|
||||
t = CompareValue(t, CMP, VALUE)
|
||||
return t
|
||||
|
||||
def VALUE(self, _parent=None):
|
||||
_context = self.Context(_parent, self._scanner, 'VALUE', [])
|
||||
_token = self._peek('"@"', 'FUNCNAME', 'STR_DQ', 'STR_SQ', context=_context)
|
||||
if _token == '"@"':
|
||||
self._scan('"@"', context=_context)
|
||||
IDENTIFIER = self._scan('IDENTIFIER', context=_context)
|
||||
return AttribValue(IDENTIFIER)
|
||||
elif _token == 'FUNCNAME':
|
||||
FUNCNAME = self._scan('FUNCNAME', context=_context)
|
||||
f = Function(FUNCNAME); args = []
|
||||
self._scan('"\\("', context=_context)
|
||||
if self._peek('"\\)"', '"@"', 'FUNCNAME', '","', 'STR_DQ', 'STR_SQ', context=_context) not in ['"\\)"', '","']:
|
||||
VALUE = self.VALUE(_context)
|
||||
args.append(VALUE)
|
||||
while self._peek('","', '"\\)"', context=_context) == '","':
|
||||
self._scan('","', context=_context)
|
||||
VALUE = self.VALUE(_context)
|
||||
args.append(VALUE)
|
||||
self._scan('"\\)"', context=_context)
|
||||
f.setParams(*args); return f
|
||||
else: # in ['STR_DQ', 'STR_SQ']
|
||||
STR = self.STR(_context)
|
||||
return LiteralValue(STR[1:len(STR)-1])
|
||||
|
||||
def CMP(self, _parent=None):
|
||||
_context = self.Context(_parent, self._scanner, 'CMP', [])
|
||||
_token = self._peek('CMP_EQ', 'CMP_NE', context=_context)
|
||||
if _token == 'CMP_EQ':
|
||||
CMP_EQ = self._scan('CMP_EQ', context=_context)
|
||||
return CMP_EQ
|
||||
else: # == 'CMP_NE'
|
||||
CMP_NE = self._scan('CMP_NE', context=_context)
|
||||
return CMP_NE
|
||||
|
||||
def STR(self, _parent=None):
|
||||
_context = self.Context(_parent, self._scanner, 'STR', [])
|
||||
_token = self._peek('STR_DQ', 'STR_SQ', context=_context)
|
||||
if _token == 'STR_DQ':
|
||||
STR_DQ = self._scan('STR_DQ', context=_context)
|
||||
return STR_DQ
|
||||
else: # == 'STR_SQ'
|
||||
STR_SQ = self._scan('STR_SQ', context=_context)
|
||||
return STR_SQ
|
||||
|
||||
|
||||
def parse(rule, text):
|
||||
P = XPathParser(XPathParserScanner(text))
|
||||
return wrap_error_reporter(P, rule)
|
||||
|
||||
if __name__ == '__main__':
|
||||
from sys import argv, stdin
|
||||
if len(argv) >= 2:
|
||||
if len(argv) >= 3:
|
||||
f = open(argv[2],'r')
|
||||
else:
|
||||
f = stdin
|
||||
print(parse(argv[1], f.read()))
|
||||
else: print ('Args: <rule> [<filename>]', file=sys.stderr)
|
||||
# End -- grammar generated by Yapps
|
||||
''')
|
||||
Reference in New Issue
Block a user